From 71bd11e42c4bdc0990b3978759ad35428a8c02e4 Mon Sep 17 00:00:00 2001 From: baihchou8787 Date: Sun, 30 Aug 2026 05:14:01 +0000 Subject: [PATCH 1/2] feat(graph-tokenizer): add strict paper reproduction workflow --- .gitattributes | 2 + .github/workflows/test_push.yml | 83 +- examples/graph_tokenizer/README.md | 304 +++ .../graph_tokenizer_trainer.py | 1737 +++++++++++++++++ examples/graph_tokenizer/paper_protocol.py | 1573 +++++++++++++++ examples/graph_tokenizer/requirements.txt | 8 + gammagl/datasets/__init__.py | 6 + gammagl/datasets/_graph_tokenizer_download.py | 228 +++ gammagl/datasets/_molecular_benchmark.py | 369 ++++ gammagl/datasets/ogbg_molhiv.py | 17 + gammagl/datasets/peptides_struct.py | 18 + gammagl/datasets/qm9.py | 23 + gammagl/models/__init__.py | 5 + gammagl/models/graph_bert.py | 309 +++ gammagl/models/graph_gte.py | 471 +++++ gammagl/models/graph_gte_pretrained.py | 129 ++ gammagl/transforms/__init__.py | 15 +- gammagl/transforms/graph_bpe.py | 276 +++ gammagl/transforms/graph_serializer.py | 568 ++++++ gammagl/transforms/graph_tokenizer.py | 379 ++++ setup.py | 22 + tests/data/test_dataset.py | 18 + .../test_graph_tokenizer_dataset_download.py | 214 ++ .../test_molecule_benchmark_datasets.py | 196 ++ tests/models/test_graph_gte_pretrained.py | 193 ++ .../test_graph_tokenizer_paper_protocol.py | 1436 ++++++++++++++ tests/models/test_graph_transformer.py | 1220 ++++++++++++ .../transforms/test_graph_bpe_cpp_install.py | 38 + tests/transforms/test_graph_tokenizer.py | 1007 ++++++++++ third_party/__init__.py | 1 + third_party/graph_bpe_cpp/__init__.py | 104 + third_party/graph_bpe_cpp/_graph_bpe.cpp | 117 ++ third_party/graph_bpe_cpp/setup.py | 30 + 33 files changed, 11113 insertions(+), 3 deletions(-) create mode 100644 .gitattributes create mode 100644 examples/graph_tokenizer/README.md create mode 100644 examples/graph_tokenizer/graph_tokenizer_trainer.py create mode 100644 examples/graph_tokenizer/paper_protocol.py create mode 100644 examples/graph_tokenizer/requirements.txt create mode 100644 gammagl/datasets/_graph_tokenizer_download.py create mode 100644 gammagl/datasets/_molecular_benchmark.py create mode 100644 gammagl/datasets/ogbg_molhiv.py create mode 100644 gammagl/datasets/peptides_struct.py create mode 100644 gammagl/datasets/qm9.py create mode 100644 gammagl/models/graph_bert.py create mode 100644 gammagl/models/graph_gte.py create mode 100644 gammagl/models/graph_gte_pretrained.py create mode 100644 gammagl/transforms/graph_bpe.py create mode 100644 gammagl/transforms/graph_serializer.py create mode 100644 gammagl/transforms/graph_tokenizer.py create mode 100644 tests/datasets/test_graph_tokenizer_dataset_download.py create mode 100644 tests/datasets/test_molecule_benchmark_datasets.py create mode 100644 tests/models/test_graph_gte_pretrained.py create mode 100644 tests/models/test_graph_tokenizer_paper_protocol.py create mode 100644 tests/models/test_graph_transformer.py create mode 100644 tests/transforms/test_graph_bpe_cpp_install.py create mode 100644 tests/transforms/test_graph_tokenizer.py create mode 100644 third_party/__init__.py create mode 100644 third_party/graph_bpe_cpp/__init__.py create mode 100644 third_party/graph_bpe_cpp/_graph_bpe.cpp create mode 100644 third_party/graph_bpe_cpp/setup.py diff --git a/.gitattributes b/.gitattributes new file mode 100644 index 000000000..6be993c81 --- /dev/null +++ b/.gitattributes @@ -0,0 +1,2 @@ +# Keep text files in the repository's canonical LF form across platforms. +* text=auto eol=lf diff --git a/.github/workflows/test_push.yml b/.github/workflows/test_push.yml index aec156114..ab01f2981 100644 --- a/.github/workflows/test_push.yml +++ b/.github/workflows/test_push.yml @@ -1,6 +1,12 @@ name: Build and Test -on: [push, pull_request] +on: + push: + pull_request: + workflow_dispatch: + schedule: + # Official GTE alignment is expensive; keep it exercised outside ordinary PRs. + - cron: '0 4 * * 0' jobs: build-and-test: @@ -22,7 +28,7 @@ jobs: python -m pip install --upgrade pip pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu pip install -r requirements.txt - pip install pybind11 ninja + pip install pybind11 ninja huggingface-hub safetensors - name: Install package run: | @@ -32,3 +38,76 @@ jobs: run: | TL_BACKEND=torch python -m compileall -q gammagl tests examples TL_BACKEND=torch python -m pytest tests/test_public_api.py tests/data tests/utils -q + + - name: Run GraphTokenizer regressions + run: | + TL_BACKEND=torch python -m pytest -q \ + tests/transforms/test_graph_tokenizer.py \ + tests/transforms/test_graph_bpe_cpp_install.py \ + tests/datasets/test_graph_tokenizer_dataset_download.py \ + tests/models/test_graph_transformer.py \ + tests/models/test_graph_tokenizer_paper_protocol.py \ + tests/models/test_graph_gte_pretrained.py::test_sha256_file_is_content_hash_and_detects_mismatch \ + tests/models/test_graph_gte_pretrained.py::test_download_rejects_checkpoint_sha256_mismatch \ + tests/models/test_graph_gte_pretrained.py::test_explicit_converter_has_full_coverage_and_parameter_equality \ + tests/models/test_graph_gte_pretrained.py::test_converter_rejects_shape_mismatch + + graph-gte-official-integration: + # The full checkpoint is roughly 610 MB. It is intentionally scheduled or + # manually dispatched, with a cache, rather than downloaded on every PR. + if: github.event_name == 'workflow_dispatch' || github.event_name == 'schedule' + runs-on: ubuntu-latest + steps: + - name: Check out repository code + uses: actions/checkout@v3 + with: + submodules: 'recursive' + + - name: Set up Python 3.10 + uses: actions/setup-python@v4 + with: + python-version: '3.10' + + - name: Restore official GTE checkpoint cache + uses: actions/cache@v3 + with: + path: ~/.cache/huggingface/hub + key: gte-multilingual-base-9bbca17d9273fd0d03d5725c7a4b0f6b45142062-f5a35a10faa54da7717870af1517c9b41e9bd8e3880bc5a8e9363d4c3c63e9b0 + + - name: Install integration dependencies + run: | + python -m pip install --upgrade pip + pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cpu + pip install -r requirements.txt + pip install pybind11 ninja huggingface-hub safetensors transformers + GAMMAGL_WITH_CUDA=0 pip install -e ".[dev]" --no-build-isolation + + - name: Download fixed official GTE checkpoint + run: | + python - <<'PY' + from huggingface_hub import snapshot_download + + snapshot_download( + repo_id="Alibaba-NLP/gte-multilingual-base", + revision="9bbca17d9273fd0d03d5725c7a4b0f6b45142062", + allow_patterns=["*.json", "*.py", "model.safetensors"], + ) + PY + + - name: Run official GTE encoder regression + run: | + GTE_CHECKPOINT_PATH="$(python - <<'PY' + from huggingface_hub import hf_hub_download + + print(hf_hub_download( + repo_id="Alibaba-NLP/gte-multilingual-base", + filename="model.safetensors", + revision="9bbca17d9273fd0d03d5725c7a4b0f6b45142062", + local_files_only=True, + )) + PY + )" + export GTE_CHECKPOINT_PATH + TL_BACKEND=torch python -m pytest -q \ + tests/models/test_graph_gte_pretrained.py::test_official_checkpoint_converter_and_hf_encoder_equivalence \ + tests/models/test_graph_gte_pretrained.py::test_official_from_pretrained_smoke_updates_encoder diff --git a/examples/graph_tokenizer/README.md b/examples/graph_tokenizer/README.md new file mode 100644 index 000000000..221aef93d --- /dev/null +++ b/examples/graph_tokenizer/README.md @@ -0,0 +1,304 @@ +# GraphTokenizer + +This example provides graph serialization, Graph BPE tokenization, masked-language-model pretraining, supervised fine-tuning, checkpointing, and multi-seed evaluation for GraphTokenizer. + +## Paper + +GraphTokenizer is described in [Graph Tokenization for Bridging Graphs and Transformers](https://openreview.net/forum?id=jCctxI1BGF) (ICLR 2026). + +The paper protocol uses frequency-guided Eulerian serialization (Feuler), fits Graph BPE on the training split only, pretrains with masked language modeling, selects the best fine-tuning checkpoint by validation performance, evaluates the test split once, and reports the mean and population standard deviation over five runs. + +The Feuler serializer is reversible for its supported simple-undirected graph +domain. Paper GTE runs load the pinned official encoder into the native TLX +GraphGTE implementation; graph-token embeddings and task heads are new +initializations. + +Feuler accepts only simple undirected graphs: self-loops and parallel edges +are rejected, and graph objects that explicitly declare `directed=True` or +`is_directed()=True` are rejected. COO may store each undirected edge once or +in symmetric form; both are canonicalized to the same undirected edge. A raw +single-direction COO has no way to encode a distinct directed-graph meaning, +so it is interpreted only as that supported undirected storage form. + +## Datasets + +The paper commands below cover: + +- QM9: joint 16-target regression, reported with MAE. +- OGBG-molhiv: binary classification, reported with OGB ROC-AUC. +- Peptides-struct: 11-target regression, reported with Average MAE. + +The GammaGL dataset classes download the official GraphTokenizer release bundle +on first use. The downloaded archive is verified before extraction using its +published SHA-256: + +``` +5c437c3c0d4b7278379c0e70d57f98148e5c815d753d8cf68e2a45952bcce459 +``` + +It is cached under `/.graph_tokenizer_release`, then copied into +each dataset's normal `raw/` directory. A local archive supplied through +`GAMMAGL_GRAPH_TOKENIZER_DATA_BUNDLE=/path/to/bundle` is also verified against +the same digest. A directory value is an explicit local-development override; +it is never downloaded from a user-provided URL. The released `data.pkl(.gz)` +files are deserialized only after the automatic remote bundle has passed this +verification. + +## Requirements + +Paper mode requires: + +- `TL_BACKEND=torch` +- PyTorch 2.1.2 with CUDA 12.1 +- DGL 2.4.0 with CUDA 12.1 +- PyTorch Geometric 2.4.0 +- TensorLayerX, NumPy, OGB, and the GammaGL package +- `huggingface-hub` and `safetensors` for the pinned native-TLX GTE checkpoint converter +- the native Graph BPE extension for the commands below + +Install these without adding paper-only packages to GammaGL's core dependency +set: + +```bash +pip install -e '.[graph-tokenizer-paper]' +# Or use examples/graph_tokenizer/requirements.txt in the pinned paper environment. +``` + +`transformers` is optional: it is used only by checkpoint/reference-equivalence +tests, never by the formal TLX GraphBERT/GraphGTE training forward path. + +Build the native Graph BPE backend from the repository root: + +```bash +export TL_BACKEND=torch +export DGLBACKEND=pytorch +python third_party/graph_bpe_cpp/setup.py build_ext --inplace +``` + +The paper models use public TensorLayerX layers and operations directly; Hugging Face Transformers is not required by GraphBERT or GraphGTE. + +## How to Run + +All paper hyperparameters are passed explicitly through `argparse`. The commands use repository-relative data, cache, and result paths and can be run directly from the GammaGL repository root. + +### QM9 + BERT + +```bash +python examples/graph_tokenizer/graph_tokenizer_trainer.py \ + --protocol paper \ + --dataset qm9 \ + --model bert \ + --serialization feuler \ + --pretrain-epoch 200 \ + --n-epoch 200 \ + --pretrain-lr 0.0001 \ + --lr 0.00001 \ + --batch-size 32 \ + --weight-decay 0.1 \ + --pretrain-warmup-ratio 0.12 \ + --finetune-warmup-ratio 0.025 \ + --pretrain-max-grad-norm 2.0 \ + --finetune-max-grad-norm 0.5 \ + --mask-prob 0.09 \ + --patience 20 \ + --max-position-embeddings 8096 \ + --data-root data \ + --bpe-backend cpp \ + --paper-cache-root cache/graph_tokenizer \ + --num-merges 2000 \ + --runs 5 \ + --seed 42 \ + --paper-amp off \ + --no-paper-tf32 \ + --output-dir logs/graph_tokenizer/qm9_bert +``` + +### QM9 + GTE + +```bash +python examples/graph_tokenizer/graph_tokenizer_trainer.py \ + --protocol paper \ + --dataset qm9 \ + --model gte \ + --serialization feuler \ + --pretrain-epoch 200 \ + --n-epoch 200 \ + --pretrain-lr 0.00005 \ + --lr 0.00001 \ + --batch-size 32 \ + --weight-decay 0.1 \ + --pretrain-warmup-ratio 0.12 \ + --finetune-warmup-ratio 0.025 \ + --pretrain-max-grad-norm 2.0 \ + --finetune-max-grad-norm 0.5 \ + --mask-prob 0.09 \ + --patience 20 \ + --max-position-embeddings 8192 \ + --data-root data \ + --bpe-backend cpp \ + --paper-cache-root cache/graph_tokenizer \ + --num-merges 2000 \ + --runs 5 \ + --seed 42 \ + --paper-amp off \ + --no-paper-tf32 \ + --output-dir logs/graph_tokenizer/qm9_gte +``` + +### OGBG-molhiv + BERT + +```bash +python examples/graph_tokenizer/graph_tokenizer_trainer.py \ + --protocol paper \ + --dataset OGBG-molhiv \ + --model bert \ + --serialization feuler \ + --pretrain-epoch 200 \ + --n-epoch 200 \ + --pretrain-lr 0.0001 \ + --lr 0.00005 \ + --batch-size 32 \ + --weight-decay 0.1 \ + --pretrain-warmup-ratio 0.12 \ + --finetune-warmup-ratio 0.025 \ + --pretrain-max-grad-norm 2.0 \ + --finetune-max-grad-norm 0.5 \ + --mask-prob 0.09 \ + --patience 20 \ + --max-position-embeddings 8096 \ + --data-root data \ + --bpe-backend cpp \ + --paper-cache-root cache/graph_tokenizer \ + --num-merges 2000 \ + --runs 5 \ + --seed 42 \ + --paper-amp off \ + --no-paper-tf32 \ + --output-dir logs/graph_tokenizer/molhiv_bert +``` + +### OGBG-molhiv + GTE + +```bash +python examples/graph_tokenizer/graph_tokenizer_trainer.py \ + --protocol paper \ + --dataset OGBG-molhiv \ + --model gte \ + --serialization feuler \ + --pretrain-epoch 200 \ + --n-epoch 200 \ + --pretrain-lr 0.00005 \ + --lr 0.00005 \ + --batch-size 32 \ + --weight-decay 0.1 \ + --pretrain-warmup-ratio 0.12 \ + --finetune-warmup-ratio 0.025 \ + --pretrain-max-grad-norm 2.0 \ + --finetune-max-grad-norm 0.5 \ + --mask-prob 0.09 \ + --patience 20 \ + --max-position-embeddings 8192 \ + --data-root data \ + --bpe-backend cpp \ + --paper-cache-root cache/graph_tokenizer \ + --num-merges 2000 \ + --runs 5 \ + --seed 42 \ + --paper-amp off \ + --no-paper-tf32 \ + --output-dir logs/graph_tokenizer/molhiv_gte +``` + +### Peptides-struct + BERT + +```bash +python examples/graph_tokenizer/graph_tokenizer_trainer.py \ + --protocol paper \ + --dataset Peptides-struct \ + --model bert \ + --serialization feuler \ + --pretrain-epoch 200 \ + --n-epoch 200 \ + --pretrain-lr 0.0001 \ + --lr 0.00001 \ + --batch-size 16 \ + --weight-decay 0.1 \ + --pretrain-warmup-ratio 0.12 \ + --finetune-warmup-ratio 0.025 \ + --pretrain-max-grad-norm 2.0 \ + --finetune-max-grad-norm 0.5 \ + --mask-prob 0.09 \ + --patience 20 \ + --max-position-embeddings 8096 \ + --data-root data \ + --bpe-backend cpp \ + --paper-cache-root cache/graph_tokenizer \ + --num-merges 2000 \ + --runs 5 \ + --seed 42 \ + --paper-amp off \ + --no-paper-tf32 \ + --output-dir logs/graph_tokenizer/peptides_struct_bert +``` + +### Peptides-struct + GTE + +```bash +python examples/graph_tokenizer/graph_tokenizer_trainer.py \ + --protocol paper \ + --dataset Peptides-struct \ + --model gte \ + --serialization feuler \ + --pretrain-epoch 200 \ + --n-epoch 200 \ + --pretrain-lr 0.0001 \ + --lr 0.00001 \ + --batch-size 16 \ + --weight-decay 0.1 \ + --pretrain-warmup-ratio 0.12 \ + --finetune-warmup-ratio 0.025 \ + --pretrain-max-grad-norm 2.0 \ + --finetune-max-grad-norm 0.5 \ + --mask-prob 0.09 \ + --patience 20 \ + --max-position-embeddings 8192 \ + --data-root data \ + --bpe-backend cpp \ + --paper-cache-root cache/graph_tokenizer \ + --num-merges 2000 \ + --runs 5 \ + --seed 42 \ + --paper-amp off \ + --no-paper-tf32 \ + --output-dir logs/graph_tokenizer/peptides_struct_gte +``` + +Before a long run, add `--preflight` to the corresponding command to validate the dataset, model, runtime, and BPE backend without training. `--resume` restores the latest phase, optimizer, scheduler, random-number-generator, and AMP scaler state from `last_state.pt`; `best.pt` remains the lightweight best-model checkpoint. + +For a small non-paper smoke test: + +```bash +python examples/graph_tokenizer/graph_tokenizer_trainer.py \ + --smoke \ + --dataset qm9 \ + --model bert \ + --data-root data \ + --bpe-backend python +``` + +## Results + +The paper reports the following five-run mean results: + +| Dataset | Encoder | Metric | Paper | GammaGL status | +| --- | --- | --- | ---: | --- | +| QM9 | BERT | raw MAE ↓ | 0.122 | revalidation required | +| QM9 | GTE | raw MAE ↓ | 0.071 | revalidation required: official weights | +| OGBG-molhiv | BERT | ROC-AUC ↑ | 82.6% | revalidation required | +| OGBG-molhiv | GTE | ROC-AUC ↑ | 87.4% | revalidation required: official weights | +| Peptides-struct | BERT | Average MAE ↓ | 0.247 | revalidation required | +| Peptides-struct | GTE | Average MAE ↓ | 0.242 | revalidation required: official weights | + +`Paper` is the value reported by the GraphTokenizer paper. GammaGL results must not be compared or described as reproduced until the strict run records the official GTE checkpoint provenance and reports QM9 in raw label units. + +Each run keeps result artifacts under `--output-dir`, including the summary JSON, per-run CSV, Markdown/LaTeX paper tables, runtime manifest, checkpoints, and paper-protocol state. JSON remains an output format only; it is not used as an input parameter configuration. diff --git a/examples/graph_tokenizer/graph_tokenizer_trainer.py b/examples/graph_tokenizer/graph_tokenizer_trainer.py new file mode 100644 index 000000000..9240cc3e1 --- /dev/null +++ b/examples/graph_tokenizer/graph_tokenizer_trainer.py @@ -0,0 +1,1737 @@ +import argparse +import copy +import csv +import hashlib +import importlib.metadata +import importlib.util +import json +import math +import os +import pickle +import platform +import random +import sys +import types +from dataclasses import dataclass +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Dict, Iterable, List, Optional, Sequence, Set, Tuple + + +@dataclass(frozen=True) +class DatasetSpec: + canonical_name: str + aliases: Tuple[str, ...] + task_type: str + output_dim: int + metric: str + paper_name: str + data_dir_names: Tuple[str, ...] = () + label_keys: Tuple[str, ...] = () + + +SUPPORTED_MODELS = ("bert", "gte") + + +DATASET_SPECS = ( + DatasetSpec( + canonical_name="qm9", + aliases=("qm9",), + task_type="regression", + output_dim=16, + metric="mae", + paper_name="QM9", + data_dir_names=("qm9",), + label_keys=( + "mu", "alpha", "homo", "lumo", "gap", "r2", "zpve", "u0", + "u298", "h298", "g298", "cv", "u0_atom", "u298_atom", + "h298_atom", "g298_atom", + ), + ), + DatasetSpec( + canonical_name="molhiv", + aliases=("molhiv", "ogbg-molhiv", "ogbg_molhiv"), + task_type="binary_classification", + output_dim=1, + metric="rocauc", + paper_name="OGBG-molhiv", + data_dir_names=("molhiv", "ogbg-molhiv", "ogbg_molhiv"), + label_keys=("label",), + ), + DatasetSpec( + canonical_name="peptides-func", + aliases=("peptides-func", "peptides_func", "p-func", "p_func"), + task_type="multi_label_classification", + output_dim=10, + metric="ap", + paper_name="Peptides-func", + data_dir_names=("peptides-func", "peptides_func", "p-func", "p_func"), + label_keys=("labels",), + ), + DatasetSpec( + canonical_name="peptides-struct", + aliases=("peptides-struct", "peptides_struct", "p-struct", "p_struct"), + task_type="multi_target_regression", + output_dim=11, + metric="average_mae", + paper_name="Peptides-struct", + data_dir_names=("peptides-struct", "peptides_struct", "p-struct", "p_struct"), + label_keys=("labels",), + ), +) + + +class SyntheticGraph: + def __init__(self, edge_index, x, edge_attr=None, y=None): + self.edge_index = edge_index + self.x = x + self.edge_attr = edge_attr + self.y = y + self.num_nodes = len(x) + + +def repo_root() -> Path: + return Path(__file__).resolve().parents[2] + + +def ensure_repo_on_path() -> Path: + root = repo_root() + root_text = str(root) + if root_text not in sys.path: + sys.path.insert(0, root_text) + return root + + +def resolve_dataset_spec(name: str) -> DatasetSpec: + normalized = name.lower().replace("_", "-") + for spec in DATASET_SPECS: + if normalized in {alias.replace("_", "-") for alias in spec.aliases}: + return spec + supported = ", ".join(spec.paper_name for spec in DATASET_SPECS) + raise ValueError(f"Unsupported dataset '{name}'. Supported datasets: {supported}.") + + +def resolve_model_name(name: str) -> str: + normalized = str(name).lower() + if normalized not in SUPPORTED_MODELS: + supported = ", ".join(SUPPORTED_MODELS) + raise ValueError(f"Unsupported model '{name}'. Supported models: {supported}.") + return normalized + + +def effective_seed(seed: Optional[int]) -> int: + return 0 if seed is None else int(seed) + + +def make_synthetic_graphs(spec: DatasetSpec) -> List[SyntheticGraph]: + targets = synthetic_targets(spec) + return [ + SyntheticGraph(edge_index=[[0, 1], [1, 2]], x=[1, 2, 3], edge_attr=[1, 2], y=targets[0]), + SyntheticGraph(edge_index=[[0, 0], [1, 2]], x=[1, 2, 4], edge_attr=[1, 3], y=targets[1]), + SyntheticGraph(edge_index=[[0, 1], [1, 0]], x=[2, 1], edge_attr=[1, 1], y=targets[2]), + ] + + +def synthetic_targets(spec: DatasetSpec) -> List[List[float]]: + if spec.metric == "rocauc": + return [[0.0], [1.0], [0.0]] + if spec.metric == "ap": + rows = [] + for index in range(3): + row = [0.0] * spec.output_dim + row[index % spec.output_dim] = 1.0 + rows.append(row) + return rows + return [[float(i + j) / 10.0 for j in range(spec.output_dim)] for i in range(3)] + + +def load_benchmark_splits(data_root, spec: DatasetSpec) -> Dict[str, List[SyntheticGraph]]: + dataset = load_gammagl_benchmark_dataset(data_root, spec) + graphs = [gammagl_graph_to_synthetic_graph(dataset[index], spec) for index in range(len(dataset))] + return select_strict_splits(graphs, dataset.get_idx_split()) + + +def load_gammagl_benchmark_dataset(data_root, spec: DatasetSpec): + ensure_repo_on_path() + dataset_cls = { + "qm9": "QM9", + "molhiv": "OGBGMolHIV", + "peptides-struct": "PeptidesStruct", + }.get(spec.canonical_name) + if dataset_cls is None: + raise ValueError( + "GammaGL GraphTokenizer benchmarks support QM9, OGBG-molhiv, " + "and Peptides-struct only.") + + from gammagl import datasets as gammagl_datasets + return getattr(gammagl_datasets, dataset_cls)(root=str(data_root)) + + +def validate_split_indices(num_samples: int, split_indices: Dict[str, List[int]]) -> None: + required = ("train", "val", "test") + if set(split_indices) != set(required): + raise ValueError(f"Split files must define exactly {required}.") + seen = {} + for split_name in required: + for index in split_indices[split_name]: + if index < 0 or index >= num_samples: + raise ValueError( + f"{split_name} split index {index} is outside [0, {num_samples}).") + if index in seen: + raise ValueError( + f"{split_name} split index {index} overlaps {seen[index]} split.") + seen[index] = split_name + + +def select_strict_splits(graphs: Sequence[SyntheticGraph], split_indices: Dict[str, List[int]]): + validate_split_indices(len(graphs), split_indices) + splits = {} + for split_name in ("train", "val", "test"): + split_graphs = [graphs[index] for index in split_indices[split_name]] + for graph, graph_id in zip(split_graphs, split_indices[split_name]): + graph.graph_tokenizer_id = int(graph_id) + splits[split_name] = split_graphs + return splits + + +def gammagl_graph_to_synthetic_graph(graph, spec: DatasetSpec) -> SyntheticGraph: + return SyntheticGraph( + edge_index=normalize_edge_index(getattr(graph, "edge_index", None)), + x=flatten_feature_ids(getattr(graph, "x", None)), + edge_attr=flatten_feature_ids(getattr(graph, "edge_attr", None)), + y=as_float_list( + getattr(graph, "y", None), + spec.output_dim, + allow_nan=spec.canonical_name == "peptides-struct", + ), + ) + + +def normalize_edge_index(edge_index) -> List[List[int]]: + if edge_index is None: + return [[], []] + values = to_list(edge_index) + if len(values) == 2 and is_sequence(values[0]) and is_sequence(values[1]): + return [[int(value) for value in values[0]], [int(value) for value in values[1]]] + src = [int(pair[0]) for pair in values] + dst = [int(pair[1]) for pair in values] + return [src, dst] + + +def flatten_feature_ids(values) -> List[int]: + if values is None: + return [] + values = to_list(values) + flattened = [] + for item in values: + if is_sequence(item): + if len(item) != 1 or is_sequence(item[0]): + raise ValueError( + "Multi-dimensional features require a dataset-specific token adapter.") + item = item[0] + flattened.append(int(item)) + return flattened + + +def as_float_list(value, output_dim: int, allow_nan: bool = False) -> List[float]: + if value is None: + raise ValueError("Sample is missing its label.") + values = to_list(value) + while is_sequence(values) and len(values) == 1 and is_sequence(values[0]): + values = values[0] + if not is_sequence(values): + values = [values] + if len(values) != output_dim: + raise ValueError( + f"Label must have exactly {output_dim} values; received {len(values)}.") + result = [float(item) for item in values] + if not allow_nan and any(math.isnan(item) for item in result): + raise ValueError("Labels cannot contain NaN values for this dataset.") + return result + + +def to_list(value): + if hasattr(value, "detach"): + value = value.detach() + if hasattr(value, "cpu"): + value = value.cpu() + if hasattr(value, "numpy"): + value = value.numpy() + if hasattr(value, "tolist"): + return value.tolist() + return value + + +def is_sequence(value: Any) -> bool: + return isinstance(value, Iterable) and not isinstance(value, (str, bytes, dict)) + + +def compute_primary_metric(spec: DatasetSpec, y_true, y_pred) -> float: + if spec.metric == "average_mae": + return average_mae(y_true, y_pred) + if spec.metric == "mae": + return mean_absolute_error(y_true, y_pred) + if spec.metric == "rocauc": + return binary_roc_auc(flatten_numeric(y_true), flatten_numeric(y_pred)) + if spec.metric == "ap": + true_rows = as_2d_float_rows(y_true) + pred_rows = as_2d_float_rows(y_pred) + if not true_rows: + return float("nan") + num_tasks = max(len(row) for row in true_rows) + scores = [] + for task_index in range(num_tasks): + task_true = [row[task_index] for row in true_rows if task_index < len(row)] + task_pred = [row[task_index] for row in pred_rows if task_index < len(row)] + if len(set(task_true)) < 2: + continue + scores.append(average_precision(task_true, task_pred)) + return sum(scores) / len(scores) if scores else float("nan") + if spec.metric in {"acc", "accuracy"}: + return accuracy_score(y_true, y_pred) + raise ValueError(f"Unsupported metric: {spec.metric}") + + +def compute_task_metrics(spec: DatasetSpec, y_true, y_pred) -> Dict[str, float]: + primary = compute_primary_metric(spec, y_true, y_pred) + if spec.metric == "rocauc": + return {"auc": primary} + if spec.metric == "ap": + return {"ap": primary} + if spec.metric == "average_mae": + return {"average_mae": primary, "mae": mean_absolute_error(y_true, y_pred)} + if spec.metric == "mae": + return {"mae": primary} + if spec.metric in {"acc", "accuracy"}: + return {"acc": primary} + raise ValueError(f"Unsupported metric: {spec.metric}") + + +def primary_metric_name(spec: DatasetSpec) -> str: + if spec.metric == "rocauc": + return "auc" + if spec.metric == "accuracy": + return "acc" + return spec.metric + + +def higher_is_better(spec: DatasetSpec) -> bool: + return spec.metric in {"rocauc", "ap", "acc", "accuracy"} + + +def metric_is_better(spec: DatasetSpec, value: float, best_value: Optional[float]) -> bool: + if math.isnan(value): + return False + if best_value is None or math.isnan(best_value): + return True + return value > best_value if higher_is_better(spec) else value < best_value + + +def select_best_epoch(spec: DatasetSpec, history: Sequence[Dict[str, Any]]) -> Optional[Dict[str, Any]]: + metric_name = primary_metric_name(spec) + best_entry = None + best_value = None + for entry in history: + value = entry.get("val", {}).get("metrics", {}).get(metric_name, float("nan")) + if metric_is_better(spec, float(value), best_value): + best_entry = entry + best_value = float(value) + return best_entry + + +def mean(values: Sequence[float]) -> float: + values = [float(value) for value in values if not math.isnan(float(value))] + return sum(values) / len(values) if values else float("nan") + + +def population_std(values: Sequence[float]) -> float: + values = [float(value) for value in values if not math.isnan(float(value))] + if not values: + return float("nan") + avg = sum(values) / len(values) + return math.sqrt(sum((value - avg) ** 2 for value in values) / len(values)) + + +def format_result_row(summary: Dict[str, Any]) -> Dict[str, Any]: + metric_name = summary["primary_metric"] + return { + "dataset": summary["dataset"], + "model": summary["model"], + "run": summary.get("run", 0), + "seed": summary.get("seed"), + "best_epoch": summary.get("best_epoch"), + f"val_{metric_name}": summary.get("best_val", {}).get("metrics", {}).get(metric_name, float("nan")), + f"test_{metric_name}": summary.get("best_test", {}).get("metrics", {}).get(metric_name, float("nan")), + "val_loss": summary.get("best_val", {}).get("loss", float("nan")), + "test_loss": summary.get("best_test", {}).get("loss", float("nan")), + } + + +def aggregate_run_summaries(spec: DatasetSpec, summaries: Sequence[Dict[str, Any]]) -> Dict[str, Any]: + metric_name = primary_metric_name(spec) + val_values = [ + float(summary.get("best_val", {}).get("metrics", {}).get(metric_name, float("nan"))) + for summary in summaries + ] + test_values = [ + float(summary.get("best_test", {}).get("metrics", {}).get(metric_name, float("nan"))) + for summary in summaries + ] + return { + "num_runs": len(summaries), + "primary_metric": metric_name, + "higher_is_better": higher_is_better(spec), + f"val_{metric_name}_mean": mean(val_values), + f"val_{metric_name}_std": population_std(val_values), + f"test_{metric_name}_mean": mean(test_values), + f"test_{metric_name}_std": population_std(test_values), + } + + +def aggregate_per_target_mae( + summaries: Sequence[Dict[str, Any]], + field_name: str = "per_target_mae") -> Dict[str, Any]: + aggregated = {} + for output_name, split_name in (("val", "best_val"), ("test", "best_test")): + target_names = [] + for summary in summaries: + for target_name in summary.get(split_name, {}).get( + field_name, {}): + if target_name not in target_names: + target_names.append(target_name) + if not target_names: + continue + aggregated[output_name] = {} + for target_name in target_names: + values = [ + float(summary.get(split_name, {}).get( + field_name, {}).get(target_name, float("nan"))) + for summary in summaries + ] + aggregated[output_name][target_name] = { + "mean": mean(values), + "std": population_std(values), + } + return aggregated + + +def format_metric_with_std(mean_value: float, std_value: float) -> str: + return f"{float(mean_value):.4f} +/- {float(std_value):.4f}" + + +def format_paper_table_row(summary: Dict[str, Any]) -> Dict[str, Any]: + aggregate = summary["aggregate"] + metric_name = aggregate["primary_metric"] + direction = "higher" if aggregate.get("higher_is_better", True) else "lower" + return { + "dataset": summary["dataset"], + "model": summary["model"], + "metric": metric_name, + "direction": direction, + "runs": aggregate.get("num_runs", 0), + "val": format_metric_with_std( + aggregate.get(f"val_{metric_name}_mean", float("nan")), + aggregate.get(f"val_{metric_name}_std", float("nan")), + ), + "test": format_metric_with_std( + aggregate.get(f"test_{metric_name}_mean", float("nan")), + aggregate.get(f"test_{metric_name}_std", float("nan")), + ), + } + + +def format_markdown_result_table(summaries: Sequence[Dict[str, Any]]) -> str: + lines = [ + "| Dataset | Model | Metric | Direction | Runs | Val | Test |", + "|---|---|---|---|---:|---:|---:|", + ] + for summary in summaries: + row = format_paper_table_row(summary) + lines.append( + f"| {row['dataset']} | {row['model']} | {row['metric']} | {row['direction']} | " + f"{row['runs']} | {row['val']} | {row['test']} |" + ) + return "\n".join(lines) + + +def format_latex_result_rows(summaries: Sequence[Dict[str, Any]]) -> str: + rows = [] + for summary in summaries: + row = format_paper_table_row(summary) + val = row["val"].replace("+/-", "$\\pm$") + test = row["test"].replace("+/-", "$\\pm$") + rows.append( + f"{row['dataset']} & {row['model']} & {row['metric']} & {row['direction']} & " + f"{row['runs']} & {val} & {test} \\" + ) + return "\n".join(rows) + + +def mean_absolute_error(y_true, y_pred) -> float: + pairs = finite_numeric_pairs(y_true, y_pred) + if not pairs: + return float("nan") + return sum(abs(true - pred) for true, pred in pairs) / len(pairs) + + +def average_mae(y_true, y_pred) -> float: + true_rows = as_2d_float_rows(y_true) + pred_rows = as_2d_float_rows(y_pred) + if not true_rows: + return float("nan") + num_tasks = max(len(row) for row in true_rows) + scores = [] + for task_index in range(num_tasks): + task_true = [row[task_index] for row in true_rows if task_index < len(row)] + task_pred = [row[task_index] for row in pred_rows if task_index < len(row)] + pairs = finite_numeric_pairs(task_true, task_pred) + if not pairs: + continue + scores.append(sum(abs(true - pred) for true, pred in pairs) / len(pairs)) + return sum(scores) / len(scores) if scores else float("nan") + + +def flatten_numeric(values) -> List[float]: + values = to_list(values) + if not is_sequence(values): + return [float(values)] + result = [] + for value in values: + if is_sequence(value): + result.extend(flatten_numeric(value)) + else: + result.append(float(value)) + return result + + +def finite_numeric_pairs(y_true, y_pred) -> List[Tuple[float, float]]: + pairs = [] + for true, pred in zip(flatten_numeric(y_true), flatten_numeric(y_pred)): + if math.isnan(true) or math.isnan(pred): + continue + pairs.append((true, pred)) + return pairs + + +def as_2d_float_rows(values) -> List[List[float]]: + values = to_list(values) + if not values: + return [] + if not is_sequence(values[0]): + return [[float(value)] for value in values] + return [[float(item) for item in row] for row in values] + + +def binary_roc_auc(y_true: Sequence[float], y_score: Sequence[float]) -> float: + positives = [score for label, score in zip(y_true, y_score) if int(label) == 1] + negatives = [score for label, score in zip(y_true, y_score) if int(label) == 0] + if not positives or not negatives: + return float("nan") + wins = 0.0 + for pos in positives: + for neg in negatives: + if pos > neg: + wins += 1.0 + elif math.isclose(pos, neg): + wins += 0.5 + return wins / (len(positives) * len(negatives)) + + +def average_precision(y_true: Sequence[float], y_score: Sequence[float]) -> float: + ranked = sorted(zip(y_score, y_true), key=lambda item: item[0], reverse=True) + total_positives = sum(1 for _, label in ranked if int(label) == 1) + if total_positives == 0: + return float("nan") + hits = 0 + precision_sum = 0.0 + for rank, (_, label) in enumerate(ranked, start=1): + if int(label) == 1: + hits += 1 + precision_sum += hits / rank + return precision_sum / total_positives + + +def accuracy_score(y_true, y_score, threshold: float = 0.5) -> float: + pairs = finite_numeric_pairs(y_true, y_score) + if not pairs: + return float("nan") + correct = 0 + for true, score in pairs: + prediction = 1 if float(score) >= threshold else 0 + if prediction == int(true): + correct += 1 + return correct / len(pairs) + + +def _ensure_local_package(package: str, package_path: Path) -> None: + module = sys.modules.setdefault(package, types.ModuleType(package)) + module.__path__ = [str(package_path)] + + +def _load_module(module_name: str, file_path: Path): + spec = importlib.util.spec_from_file_location(module_name, file_path) + module = importlib.util.module_from_spec(spec) + sys.modules[module_name] = module + spec.loader.exec_module(module) + return module + + +def load_tokenizer_classes(): + root = ensure_repo_on_path() + _ensure_local_package("gammagl", root / "gammagl") + _ensure_local_package("gammagl.transforms", root / "gammagl" / "transforms") + tokenizer_module = _load_module( + "gammagl.transforms.graph_tokenizer", + root / "gammagl" / "transforms" / "graph_tokenizer.py", + ) + bpe_module = _load_module( + "gammagl.transforms.graph_bpe", + root / "gammagl" / "transforms" / "graph_bpe.py", + ) + return tokenizer_module.GraphTokenizer, bpe_module.GraphBPE + + +def load_model_class(model_name: str): + model_name = resolve_model_name(model_name) + root = ensure_repo_on_path() + _ensure_local_package("gammagl", root / "gammagl") + _ensure_local_package("gammagl.models", root / "gammagl" / "models") + graph_bert = _load_module("gammagl.models.graph_bert", root / "gammagl" / "models" / "graph_bert.py") + if model_name == "bert": + return graph_bert.GraphBERT + graph_gte = _load_module("gammagl.models.graph_gte", root / "gammagl" / "models" / "graph_gte.py") + return graph_gte.GraphGTE + + +def model_kwargs(args, spec: DatasetSpec): + if args.smoke: + return { + "vocab_size": args.vocab_size, + "output_dim": spec.output_dim, + "hidden_size": args.hidden_size or 24, + "num_hidden_layers": args.num_hidden_layers or 1, + "num_attention_heads": args.num_attention_heads or 4, + "intermediate_size": args.intermediate_size or 48, + "max_position_embeddings": args.max_length, + } + kwargs = { + "vocab_size": args.vocab_size, + "output_dim": spec.output_dim, + "max_position_embeddings": args.max_length, + } + optional_values = { + "hidden_size": args.hidden_size, + "num_hidden_layers": args.num_hidden_layers, + "num_attention_heads": args.num_attention_heads, + "intermediate_size": args.intermediate_size, + } + kwargs.update({key: value for key, value in optional_values.items() if value is not None}) + return kwargs + + +def load_graphs(args, spec: DatasetSpec): + if args.smoke: + return make_synthetic_graphs(spec) + splits = load_benchmark_splits(args.data_root, spec) + return splits["train"] + splits["val"] + splits["test"] + + +def fit_tokenizer(args, train_graphs: Sequence[SyntheticGraph]): + GraphTokenizer, GraphBPE = load_tokenizer_classes() + from gammagl.transforms.graph_tokenizer import GraphTokenizerSpecialTokens + from gammagl.transforms.graph_serializer import ( + EulerianSerializer, + FrequencyGuidedEulerianSerializer, + ) + + serialization = getattr(args, "serialization", "feuler") + serializer_cls = { + "feuler": FrequencyGuidedEulerianSerializer, + "eulerian": EulerianSerializer, + }.get(serialization) + if serializer_cls is None: + raise ValueError("serialization must be 'feuler' or 'eulerian'.") + special_tokens = GraphTokenizerSpecialTokens() + if getattr(args, "protocol", None) == "paper" and getattr(args, "model", None) == "gte": + # The pinned GTE checkpoint fixes padding_idx=1. Keep the eight + # special IDs contiguous and distinct while preserving BERT's PAD=0. + special_tokens = GraphTokenizerSpecialTokens( + pad_token_id=1, + unk_token_id=0, + ) + tokenizer = GraphTokenizer( + serializer=serializer_cls(), + bpe=GraphBPE( + num_merges=args.num_merges, + min_frequency=args.min_frequency, + backend=args.bpe_backend, + ), + special_tokens=special_tokens, + ) + if getattr(args, "protocol", None) == "paper": + cache_root = Path( + getattr(args, "paper_cache_root", None) + or (Path(args.data_root) / ".graph_tokenizer_cache") + ) + digest = hashlib.sha256() + cache_config = { + "cache_version": 2, + "dataset": str(args.dataset), + "model": str(getattr(args, "model", None)), + "serialization": serialization, + "num_merges": int(args.num_merges), + "min_frequency": int(args.min_frequency), + "bpe_backend": str(args.bpe_backend), + "num_train_graphs": len(train_graphs), + } + digest.update(json.dumps( + cache_config, sort_keys=True, separators=(",", ":") + ).encode("utf-8")) + for graph in train_graphs: + digest.update(pickle.dumps(( + to_list(graph.edge_index), + to_list(graph.x), + to_list(graph.edge_attr), + int(graph.num_nodes), + ), protocol=4)) + cache_key = digest.hexdigest() + cache_path = cache_root / str(args.dataset) / cache_key / "tokenizer.pkl" + if cache_path.is_file(): + try: + with cache_path.open("rb") as handle: + cached = pickle.load(handle) + if ( + cached.get("cache_version") == 2 + and cached.get("cache_key") == cache_key): + tokenizer = cached["tokenizer"] + tokenizer._cache_status = "hit" + tokenizer._cache_key = cache_key + tokenizer._cache_path = str(cache_path) + return tokenizer + except (OSError, EOFError, AttributeError, KeyError, pickle.UnpicklingError): + pass + tokenizer.fit( + train_graphs, + graph_ids=[ + getattr(graph, "graph_tokenizer_id", index) + for index, graph in enumerate(train_graphs) + ], + ) + if getattr(args, "protocol", None) == "paper": + tokenizer._cache_status = "miss" + tokenizer._cache_key = cache_key + tokenizer._cache_path = str(cache_path) + cache_path.parent.mkdir(parents=True, exist_ok=True) + temporary = cache_path.with_name( + f".{cache_path.name}.tmp-{os.getpid()}") + with temporary.open("wb") as handle: + pickle.dump({ + "cache_version": 2, + "cache_key": cache_key, + "tokenizer": tokenizer, + }, handle, protocol=pickle.HIGHEST_PROTOCOL) + os.replace(temporary, cache_path) + return tokenizer + + +def encode_graph_batch(graphs: Sequence[SyntheticGraph], tokenizer, max_length: int) -> Dict[str, List[List[float]]]: + encoded = [tokenizer.encode_graph(graph).input_ids for graph in graphs] + input_ids, attention_mask = tokenizer.pad_token_sequences( + encoded, max_length=max_length) + return { + "input_ids": input_ids, + "attention_mask": attention_mask, + "labels": [graph.y for graph in graphs], + } + + +def encode_graph_splits( + splits: Dict[str, List[SyntheticGraph]], + tokenizer, + max_length: int, +) -> Dict[str, Dict[str, List[List[float]]]]: + return { + split_name: encode_graph_batch(graphs, tokenizer, max_length=max_length) + for split_name, graphs in splits.items() + } + + +def iter_token_batches(encoded_split: Dict[str, List[List[float]]], batch_size: int): + total = len(encoded_split["input_ids"]) + batch_size = max(1, int(batch_size)) + for start in range(0, total, batch_size): + end = min(start + batch_size, total) + yield { + "input_ids": encoded_split["input_ids"][start:end], + "attention_mask": encoded_split["attention_mask"][start:end], + "labels": encoded_split["labels"][start:end], + } + + +def build_mlm_pretrain_split( + encoded_split: Dict[str, List[List[float]]], + tokenizer, + mask_prob: float, + seed: int, +) -> Dict[str, List[List[int]]]: + masked_ids, mlm_labels = tokenizer.mask_token_sequences( + encoded_split["input_ids"], + mask_prob=mask_prob, + seed=seed, + ) + return { + "input_ids": masked_ids, + "attention_mask": encoded_split["attention_mask"], + "mlm_labels": mlm_labels, + } + + +def count_mlm_labels(mlm_labels: Sequence[Sequence[int]]) -> int: + return sum(1 for row in mlm_labels for value in row if int(value) != -100) + + +def iter_mlm_batches(encoded_split: Dict[str, List[List[int]]], batch_size: int): + total = len(encoded_split["input_ids"]) + batch_size = max(1, int(batch_size)) + for start in range(0, total, batch_size): + end = min(start + batch_size, total) + yield { + "input_ids": encoded_split["input_ids"][start:end], + "attention_mask": encoded_split["attention_mask"][start:end], + "mlm_labels": encoded_split["mlm_labels"][start:end], + } + + +def tensor_to_float(value) -> float: + value = to_list(value) + if is_sequence(value): + flat = flatten_numeric(value) + return float(flat[0]) if flat else float("nan") + return float(value) + + +def logits_to_predictions(spec: DatasetSpec, logits) -> List[List[float]]: + rows = as_2d_float_rows(logits) + if spec.metric in {"rocauc", "ap"}: + return [[1.0 / (1.0 + math.exp(-value)) for value in row] for row in rows] + return rows + + +def supervised_loss(logits, labels, spec: DatasetSpec): + import tensorlayerx as tlx + + if spec.metric in {"mae", "average_mae"}: + return tlx.losses.mean_squared_error(logits, labels) + return tlx.losses.sigmoid_cross_entropy(target=labels, output=logits) + + +def masked_language_modeling_loss(mlm_logits, mlm_labels): + try: + import torch + + if isinstance(mlm_logits, torch.Tensor): + vocab_size = int(mlm_logits.shape[-1]) + return torch.nn.functional.cross_entropy( + mlm_logits.reshape(-1, vocab_size), + mlm_labels.reshape(-1).long(), + ignore_index=-100, + ) + except ImportError: + pass + raise RuntimeError("MLM pretraining currently requires TensorLayerX with the PyTorch backend.") + + +class GraphTokenMLMLoss: + def __new__(cls, model): + from tensorlayerx.model import WithLoss + + class _Loss(WithLoss): + def __init__(self, backbone): + super().__init__(backbone=backbone, loss_fn=None) + + def forward(self, input_ids, attention_mask, mlm_labels): + outputs = self.backbone_network(input_ids, attention_mask=attention_mask) + return masked_language_modeling_loss(outputs["mlm_logits"], mlm_labels) + + return _Loss(model) + + +class GraphTokenSupervisedLoss: + def __new__(cls, model, spec: DatasetSpec): + import tensorlayerx as tlx + from tensorlayerx.model import WithLoss + + class _Loss(WithLoss): + def __init__(self, backbone, dataset_spec): + super().__init__(backbone=backbone, loss_fn=None) + self.dataset_spec = dataset_spec + + def forward(self, input_ids, attention_mask, labels): + outputs = self.backbone_network(input_ids, attention_mask=attention_mask) + return supervised_loss(outputs["logits"], labels, self.dataset_spec) + + return _Loss(model, spec) + + +def train_one_epoch(model, train_one_step, encoded_train, spec: DatasetSpec, args) -> float: + import tensorlayerx as tlx + + model.set_train() + losses = [] + for batch in iter_token_batches(encoded_train, args.batch_size): + if not batch["input_ids"]: + continue + loss = train_one_step( + tlx.convert_to_tensor(batch["input_ids"], dtype=tlx.int64), + tlx.convert_to_tensor(batch["attention_mask"], dtype=tlx.int64), + tlx.convert_to_tensor(batch["labels"], dtype=tlx.float32), + ) + losses.append(tensor_to_float(loss)) + return sum(losses) / len(losses) if losses else float("nan") + + +def evaluate_split(model, encoded_split, spec: DatasetSpec, args) -> Dict[str, Any]: + import tensorlayerx as tlx + + model.set_eval() + losses = [] + y_true = [] + y_pred = [] + for batch in iter_token_batches(encoded_split, args.batch_size): + if not batch["input_ids"]: + continue + labels = tlx.convert_to_tensor(batch["labels"], dtype=tlx.float32) + outputs = model( + tlx.convert_to_tensor(batch["input_ids"], dtype=tlx.int64), + attention_mask=tlx.convert_to_tensor(batch["attention_mask"], dtype=tlx.int64), + ) + loss = supervised_loss(outputs["logits"], labels, spec) + losses.append(tensor_to_float(loss)) + y_true.extend(batch["labels"]) + y_pred.extend(logits_to_predictions(spec, to_list(outputs["logits"]))) + metrics = compute_task_metrics(spec, y_true, y_pred) if y_true else {} + return { + "loss": sum(losses) / len(losses) if losses else float("nan"), + "metrics": metrics, + "num_graphs": len(encoded_split["labels"]), + } + + +def run_mlm_pretrain(model, encoded_train, tokenizer, args) -> Dict[str, Any]: + if args.pretrain_epoch <= 0 or not encoded_train["input_ids"]: + return {"epochs_ran": 0, "final_loss": float("nan"), "history": []} + + import tensorlayerx as tlx + from tensorlayerx.model import TrainOneStep + + special_tokens = tokenizer.special_tokens + special_token_ids = { + special_tokens.pad_token_id, + special_tokens.cls_token_id, + special_tokens.sep_token_id, + special_tokens.component_sep_token_id, + } + optimizer = tlx.optimizers.Adam(lr=args.pretrain_lr, weight_decay=args.weight_decay) + train_one_step = TrainOneStep(GraphTokenMLMLoss(model), optimizer, model.trainable_weights) + + history = [] + for epoch in range(1, args.pretrain_epoch + 1): + model.set_train() + masked_train = build_mlm_pretrain_split( + encoded_train, + tokenizer=tokenizer, + mask_prob=args.mask_prob, + seed=effective_seed(args.seed) + epoch, + ) + num_mlm_labels = count_mlm_labels(masked_train["mlm_labels"]) + losses = [] + if num_mlm_labels > 0: + for batch in iter_mlm_batches(masked_train, args.batch_size): + loss = train_one_step( + tlx.convert_to_tensor(batch["input_ids"], dtype=tlx.int64), + tlx.convert_to_tensor(batch["attention_mask"], dtype=tlx.int64), + tlx.convert_to_tensor(batch["mlm_labels"], dtype=tlx.int64), + ) + losses.append(tensor_to_float(loss)) + history.append( + { + "epoch": epoch, + "loss": sum(losses) / len(losses) if losses else float("nan"), + "num_mlm_labels": num_mlm_labels, + } + ) + + return { + "epochs_ran": len(history), + "final_loss": history[-1]["loss"] if history else float("nan"), + "history": history, + } + + +def checkpoint_path(args, spec: DatasetSpec, run_index: int) -> Path: + return ( + Path(args.output_dir) + / "checkpoints" + / spec.canonical_name + / args.model + / f"run_{run_index}_best.npz" + ) + + +def save_model_checkpoint(model, path: Path) -> str: + path.parent.mkdir(parents=True, exist_ok=True) + model.save_weights(str(path), format="npz_dict") + return str(path) + + +def run_train_val_test(args): + spec = resolve_dataset_spec(args.dataset) + random.seed(effective_seed(args.seed)) + splits = load_benchmark_splits(args.data_root, spec) + train_graphs = splits["train"] + if not train_graphs: + raise ValueError("GraphTokenizer training split cannot be empty.") + tokenizer = fit_tokenizer(args, train_graphs) + encoded_splits = encode_graph_splits(splits, tokenizer, max_length=args.max_length) + + import tensorlayerx as tlx + from tensorlayerx.model import TrainOneStep + + Model = load_model_class(args.model) + model = Model(**model_kwargs(args, spec)) + pretrain_summary = run_mlm_pretrain(model, encoded_splits["train"], tokenizer, args) + optimizer = tlx.optimizers.Adam(lr=args.lr, weight_decay=args.weight_decay) + loss_func = GraphTokenSupervisedLoss(model, spec) + train_one_step = TrainOneStep(loss_func, optimizer, model.trainable_weights) + + history = [] + best_value = None + best_epoch = None + best_checkpoint_path = None + best_model_state = None + stale_epochs = 0 + metric_name = primary_metric_name(spec) + run_index = int(getattr(args, "run", 0)) + for epoch in range(1, args.n_epoch + 1): + train_loss = train_one_epoch(model, train_one_step, encoded_splits["train"], spec, args) + val_result = evaluate_split(model, encoded_splits["val"], spec, args) + entry = { + "epoch": epoch, + "train_loss": train_loss, + "val": val_result, + } + history.append(entry) + val_metric = val_result["metrics"].get(metric_name, float("nan")) + if metric_is_better(spec, float(val_metric), best_value): + best_value = float(val_metric) + best_epoch = epoch + stale_epochs = 0 + best_model_state = copy.deepcopy(model.state_dict()) + if args.save_checkpoint: + best_checkpoint_path = save_model_checkpoint(model, checkpoint_path(args, spec, run_index)) + else: + stale_epochs += 1 + if args.patience > 0 and stale_epochs >= args.patience: + break + + best = select_best_epoch(spec, history) or (history[-1] if history else None) + if best_model_state is None: + raise RuntimeError("No validation checkpoint was selected.") + model.load_state_dict(best_model_state) + best_test = evaluate_split(model, encoded_splits["test"], spec, args) + return { + "dataset": spec.canonical_name, + "model": args.model, + "run": run_index, + "seed": effective_seed(args.seed), + "metric": spec.metric, + "primary_metric": metric_name, + "higher_is_better": higher_is_better(spec), + "train": len(splits["train"]), + "val": len(splits["val"]), + "test": len(splits["test"]), + "epochs_ran": len(history), + "best_epoch": best_epoch if best_epoch is not None else (best["epoch"] if best else None), + "pretrain": pretrain_summary, + "checkpoint_path": best_checkpoint_path, + "best_val": best["val"] if best else {}, + "best_test": best_test, + "history": history, + } + + +def clone_args_for_run(args, run_index: int): + run_args = argparse.Namespace(**vars(args)) + run_args.run = run_index + run_args.seed = effective_seed(args.seed) + int(run_index) + return run_args + + +def parse_csv_values(value) -> List[str]: + if value is None: + return [] + if isinstance(value, (list, tuple)): + return [str(item).strip() for item in value if str(item).strip()] + return [item.strip() for item in str(value).split(",") if item.strip()] + + +def experiment_matrix_items(args) -> List[Tuple[str, str]]: + datasets = parse_csv_values(getattr(args, "datasets", None)) or [args.dataset] + models = parse_csv_values(getattr(args, "models", None)) or [args.model] + return [(dataset, model) for dataset in datasets for model in models] + + +def clone_args_for_experiment(args, dataset: str, model: str): + experiment_args = argparse.Namespace(**vars(args)) + experiment_args.dataset = dataset + experiment_args.model = model + return experiment_args + + + +def serialize_arg_value(value): + if value is None or isinstance(value, (str, int, float, bool)): + return value + if isinstance(value, Path): + return str(value) + if isinstance(value, (list, tuple, set)): + return [serialize_arg_value(item) for item in value] + if isinstance(value, dict): + return {str(key): serialize_arg_value(item) for key, item in value.items()} + return str(value) + + +def serialize_args(args) -> Dict[str, Any]: + arg_dict = vars(args) if hasattr(args, "__dict__") else {} + if arg_dict: + items = arg_dict.items() + else: + items = ( + (name, getattr(args, name)) + for name in dir(args) + if not name.startswith("_") and not callable(getattr(args, name)) + ) + return {name: serialize_arg_value(value) for name, value in sorted(items)} + + +def package_versions(names: Sequence[str]) -> Dict[str, Optional[str]]: + versions = {} + for name in names: + try: + versions[name] = importlib.metadata.version(name) + except importlib.metadata.PackageNotFoundError: + versions[name] = None + return versions + + +def git_metadata() -> Dict[str, Any]: + return { + "status": "disabled", + "reason": "local-only run; git commands are disabled", + } + + +def build_runtime_manifest(args, timestamp: Optional[str] = None) -> Dict[str, Any]: + return { + "timestamp": timestamp or datetime.now(timezone.utc).isoformat(), + "args": serialize_args(args), + "python": { + "version": sys.version, + "executable": sys.executable, + "platform": platform.platform(), + }, + "packages": package_versions(("numpy", "torch", "torchvision", "tensorlayerx", "gammagl")), + "git": git_metadata(), + } + + + +def parse_args(argv=None): + parser = build_arg_parser() + args = parser.parse_args(argv) + if args.model is None: + args.model = "gte" if args.protocol == "paper" else "bert" + return args + + + +def preflight_dataset_status(data_root, dataset: str) -> Dict[str, Any]: + spec = resolve_dataset_spec(dataset) + status: Dict[str, Any] = { + "dataset": spec.canonical_name, + "paper_name": spec.paper_name, + "status": "ok", + } + gammagl_dataset = load_gammagl_benchmark_dataset(data_root, spec) + if gammagl_dataset is None: + raise FileNotFoundError( + f"GammaGL dataset loader could not materialize {spec.paper_name}.") + split_indices = gammagl_dataset.get_idx_split() + validate_split_indices(len(gammagl_dataset), split_indices) + status.update( + { + "dataset_dir": str(gammagl_dataset.root_dir), + "storage_dir": str(gammagl_dataset.raw_dir), + "num_samples": len(gammagl_dataset), + "split_sizes": {split_name: len(indices) for split_name, indices in split_indices.items()}, + "prepared_before_training": True, + } + ) + return status + + +def prepare_paper_datasets(data_root, datasets: Sequence[str]) -> List[Dict[str, Any]]: + """Single-process installation of immutable paper data before training.""" + ensure_repo_on_path() + from gammagl.datasets._graph_tokenizer_download import materialize_paper_dataset + + prepared = [] + for dataset in datasets: + spec = resolve_dataset_spec(dataset) + raw_dir = Path(data_root) / spec.canonical_name / "raw" + materialize_paper_dataset( + dataset_name=spec.canonical_name, + aliases=spec.aliases, + raw_dir=raw_dir, + cache_root=Path(data_root), + allow_download=True, + ) + prepared.append(preflight_dataset_status(data_root, dataset)) + return prepared + + +def preflight_model_status(model: str) -> Dict[str, Any]: + model_name = resolve_model_name(model) + filename = "graph_bert.py" if model_name == "bert" else "graph_gte.py" + model_path = repo_root() / "gammagl" / "models" / filename + if not model_path.exists(): + raise FileNotFoundError(f"Model file not found: {model_path}") + status = {"model": model_name, "status": "ok", "file": str(model_path)} + if model_name == "gte": + status.update({ + "reproduction_ready": True, + "pretrained_model_id": "Alibaba-NLP/gte-multilingual-base", + "pretrained_revision": "9bbca17d9273fd0d03d5725c7a4b0f6b45142062", + }) + return status + + +def preflight_bpe_backend_status(backend: str) -> Dict[str, Any]: + ensure_repo_on_path() + if backend not in {"python", "auto", "cpp"}: + raise ValueError("bpe_backend must be one of: python, auto, cpp.") + status = {"backend": backend, "status": "ok"} + if backend == "cpp": + from third_party import graph_bpe_cpp + + if not graph_bpe_cpp.is_available(): + raise ImportError("GraphBPE backend='cpp' requires the optional graph_bpe_cpp native extension.") + status["native_available"] = True + elif backend == "auto": + from third_party import graph_bpe_cpp + + status["native_available"] = graph_bpe_cpp.is_available() + return status + + +def preflight_check(args) -> Dict[str, Any]: + datasets = parse_csv_values(getattr(args, "datasets", None)) or [args.dataset] + models = parse_csv_values(getattr(args, "models", None)) + if not models: + models = [args.model or "gte"] + report: Dict[str, Any] = { + "status": "ok", + "data_root": str(getattr(args, "data_root", "data")), + "datasets": [], + "models": [], + "bpe_backend": None, + "paper_runtime": None, + "errors": [], + } + if getattr(args, "protocol", "gammagl") == "paper": + try: + report["paper_runtime"] = ( + load_paper_protocol_module().paper_runtime_status()) + except Exception as error: + report["paper_runtime"] = { + "status": "failed", + "error_type": type(error).__name__, + "error": str(error), + } + report["errors"].append(str(error)) + + for dataset in datasets: + try: + if getattr(args, "protocol", "gammagl") == "paper": + protocol = load_paper_protocol_module() + protocol.validate_paper_args(args, resolve_dataset_spec(dataset)) + report["datasets"].append(preflight_dataset_status(args.data_root, dataset)) + except Exception as error: + report["datasets"].append( + { + "dataset": str(dataset), + "status": "failed", + "error_type": type(error).__name__, + "error": str(error), + } + ) + report["errors"].append(str(error)) + + for model in models: + try: + report["models"].append(preflight_model_status(model)) + except Exception as error: + report["models"].append( + { + "model": str(model), + "status": "failed", + "error_type": type(error).__name__, + "error": str(error), + } + ) + report["errors"].append(str(error)) + + try: + report["bpe_backend"] = preflight_bpe_backend_status(args.bpe_backend) + except Exception as error: + report["bpe_backend"] = { + "backend": str(getattr(args, "bpe_backend", None)), + "status": "failed", + "error_type": type(error).__name__, + "error": str(error), + } + report["errors"].append(str(error)) + + report["num_errors"] = len(report["errors"]) + if report["num_errors"]: + report["status"] = "failed" + return report + +def experiment_summary_path(args, dataset: str, model: str) -> Path: + spec = resolve_dataset_spec(dataset) + return Path(args.output_dir) / f"{spec.canonical_name}_{model}_summary.json" + + +def load_existing_experiment_summary(args, dataset: str, model: str) -> Optional[Dict[str, Any]]: + path = experiment_summary_path(args, dataset, model) + if not path.exists(): + return None + with path.open("r", encoding="utf-8") as handle: + summary = json.load(handle) + summary.setdefault("status", "skipped_existing") + summary.setdefault("outputs", {})["json"] = str(path) + return summary + + +def format_failed_experiment_summary(dataset: str, model: str, error: Exception) -> Dict[str, Any]: + return { + "dataset": dataset, + "model": model, + "status": "failed", + "error_type": type(error).__name__, + "error": str(error), + "runs": [], + "rows": [], + "aggregate": { + "num_runs": 0, + "primary_metric": None, + "higher_is_better": None, + }, + } + + +def write_experiment_outputs(summary: Dict[str, Any], args) -> Dict[str, str]: + output_dir = Path(args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + dataset = summary["dataset"] + model = summary["model"] + json_path = output_dir / f"{dataset}_{model}_summary.json" + csv_path = output_dir / f"{dataset}_{model}_runs.csv" + markdown_path = output_dir / f"{dataset}_{model}_paper_table.md" + latex_path = output_dir / f"{dataset}_{model}_paper_table.tex" + manifest_path = output_dir / f"{dataset}_{model}_manifest.json" + + with json_path.open("w", encoding="utf-8") as handle: + json.dump(summary, handle, indent=2) + + rows = summary.get("rows", []) + if rows: + fieldnames = list(rows[0].keys()) + with csv_path.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fieldnames) + writer.writeheader() + writer.writerows(rows) + + table_summaries = [summary] + markdown_path.write_text(format_markdown_result_table(table_summaries), encoding="utf-8") + latex_path.write_text(format_latex_result_rows(table_summaries), encoding="utf-8") + if "manifest" in summary: + with manifest_path.open("w", encoding="utf-8") as handle: + json.dump(summary["manifest"], handle, indent=2) + return { + "json": str(json_path), + "csv": str(csv_path), + "markdown": str(markdown_path), + "latex": str(latex_path), + "manifest": str(manifest_path), + } + + +def run_experiments(args): + spec = resolve_dataset_spec(args.dataset) + runs = [ + run_train_val_test(clone_args_for_run(args, run_index)) + for run_index in range(max(1, int(args.runs or 1))) + ] + rows = [format_result_row(summary) for summary in runs] + summary = { + "dataset": spec.canonical_name, + "model": args.model, + "status": "completed", + "runs": runs, + "rows": rows, + "aggregate": aggregate_run_summaries(spec, runs), + "manifest": build_runtime_manifest(args), + } + summary["outputs"] = write_experiment_outputs(summary, args) + return summary + + +def load_paper_protocol_module(): + module_name = "graph_tokenizer_paper_protocol" + if module_name in sys.modules: + return sys.modules[module_name] + path = Path(__file__).with_name("paper_protocol.py") + module_spec = importlib.util.spec_from_file_location(module_name, path) + module = importlib.util.module_from_spec(module_spec) + sys.modules[module_name] = module + module_spec.loader.exec_module(module) + return module + + +def run_paper_experiments(args): + spec = resolve_dataset_spec(args.dataset) + protocol = load_paper_protocol_module() + protocol.validate_paper_args(args, spec) + paper_summary = protocol.run_paper_experiment( + args, + spec, + load_benchmark_splits(args.data_root, spec), + fit_tokenizer, + ) + primary_metric = "rocauc" if spec.canonical_name == "molhiv" else spec.metric + runs = [] + for run in paper_summary["runs"]: + runs.append({ + "dataset": spec.canonical_name, + "model": paper_summary["model"], + "run": run["run"], + "seed": run["seed"], + "best_epoch": run["best_epoch"], + "best_val": { + "loss": run["best_val"]["loss"], + "metrics": {primary_metric: run["best_val"]["metric"]}, + "metric_space": run["best_val"].get("metric_space"), + "per_target_mae": run["best_val"].get( + "per_target_mae", {}), + "per_target_mae_raw": run["best_val"].get( + "per_target_mae_raw", {}), + }, + "best_test": { + "loss": run["best_test"]["loss"], + "metrics": {primary_metric: run["best_test"]["metric"]}, + "metric_space": run["best_test"].get("metric_space"), + "per_target_mae": run["best_test"].get( + "per_target_mae", {}), + "per_target_mae_raw": run["best_test"].get( + "per_target_mae_raw", {}), + }, + "checkpoint_path": run["checkpoint_path"], + "history": run["history"], + "test_evaluations": run["test_evaluations"], + "model_manifest": run["model"], + "loss_space": run["loss_space"], + "metric_space": run["metric_space"], + "target_normalization": run["target_normalization"], + }) + values = [run["best_test"]["metrics"][primary_metric] for run in runs] + val_values = [run["best_val"]["metrics"][primary_metric] for run in runs] + summary = { + "dataset": spec.canonical_name, + "model": paper_summary["model"], + "protocol": "paper", + "status": "completed", + "runs": runs, + "rows": [format_result_row({**run, "primary_metric": primary_metric}) for run in runs], + "aggregate": { + "num_runs": len(runs), + "primary_metric": primary_metric, + "higher_is_better": spec.canonical_name == "molhiv", + f"val_{primary_metric}_mean": mean(val_values), + f"val_{primary_metric}_std": population_std(val_values), + f"test_{primary_metric}_mean": mean(values), + f"test_{primary_metric}_std": population_std(values), + "per_target_mae": aggregate_per_target_mae(runs), + "per_target_mae_raw": aggregate_per_target_mae( + runs, field_name="per_target_mae_raw"), + }, + "paper": {key: value for key, value in paper_summary.items() if key != "runs"}, + "manifest": { + **build_runtime_manifest(args), + "loss_space": paper_summary["loss_space"], + "metric_space": paper_summary["metric_space"], + "target_normalization": paper_summary["target_normalization"], + }, + } + summary["outputs"] = write_experiment_outputs(summary, args) + return summary + + +def write_experiment_matrix_outputs(summary: Dict[str, Any], args) -> Dict[str, str]: + output_dir = Path(args.output_dir) + output_dir.mkdir(parents=True, exist_ok=True) + json_path = output_dir / "graph_tokenizer_all_summary.json" + csv_path = output_dir / "graph_tokenizer_all_runs.csv" + markdown_path = output_dir / "graph_tokenizer_all_paper_table.md" + latex_path = output_dir / "graph_tokenizer_all_paper_table.tex" + manifest_path = output_dir / "graph_tokenizer_all_manifest.json" + + with json_path.open("w", encoding="utf-8") as handle: + json.dump(summary, handle, indent=2) + + all_rows = [] + for experiment in summary.get("experiments", []): + all_rows.extend(experiment.get("rows", [])) + if all_rows: + fieldnames = sorted({key for row in all_rows for key in row.keys()}) + with csv_path.open("w", newline="", encoding="utf-8") as handle: + writer = csv.DictWriter(handle, fieldnames=fieldnames) + writer.writeheader() + writer.writerows(all_rows) + + experiments = summary.get("experiments", []) + markdown_path.write_text(format_markdown_result_table(experiments), encoding="utf-8") + latex_path.write_text(format_latex_result_rows(experiments), encoding="utf-8") + if "manifest" in summary: + with manifest_path.open("w", encoding="utf-8") as handle: + json.dump(summary["manifest"], handle, indent=2) + return { + "json": str(json_path), + "csv": str(csv_path), + "markdown": str(markdown_path), + "latex": str(latex_path), + "manifest": str(manifest_path), + } + + +def run_experiment_matrix(args): + experiments = [] + failures = [] + for dataset, model in experiment_matrix_items(args): + if getattr(args, "resume", False): + existing = load_existing_experiment_summary(args, dataset, model) + if existing is not None: + experiments.append(existing) + continue + try: + experiments.append(run_experiments(clone_args_for_experiment(args, dataset, model))) + except Exception as error: + failure = format_failed_experiment_summary(dataset, model, error) + failures.append(failure) + experiments.append(failure) + if not getattr(args, "keep_going", False): + summary = { + "datasets": parse_csv_values(getattr(args, "datasets", None)) or [args.dataset], + "models": parse_csv_values(getattr(args, "models", None)) or [args.model], + "num_experiments": len(experiments), + "num_failures": len(failures), + "experiments": experiments, + "failures": failures, + "manifest": build_runtime_manifest(args), + } + summary["outputs"] = write_experiment_matrix_outputs(summary, args) + raise + summary = { + "datasets": parse_csv_values(getattr(args, "datasets", None)) or [args.dataset], + "models": parse_csv_values(getattr(args, "models", None)) or [args.model], + "num_experiments": len(experiments), + "num_failures": len(failures), + "experiments": experiments, + "failures": failures, + "manifest": build_runtime_manifest(args), + } + summary["outputs"] = write_experiment_matrix_outputs(summary, args) + return summary + + +def run_smoke(args): + spec = resolve_dataset_spec(args.dataset) + graphs = load_graphs(args, spec) + tokenizer = fit_tokenizer(args, graphs) + encoded = [tokenizer.encode_graph(graph).input_ids for graph in graphs] + mlm_batch = tokenizer.build_mlm_batch_from_token_sequences( + encoded, + max_length=args.max_length, + mask_prob=args.mask_prob, + seed=effective_seed(args.seed), + ) + masked_ids = mlm_batch.input_ids + attention_mask = mlm_batch.attention_mask + mlm_labels = mlm_batch.labels + + import tensorlayerx as tlx + + Model = load_model_class(args.model) + model = Model(**model_kwargs(args, spec)) + outputs = model( + tlx.convert_to_tensor(masked_ids, dtype=tlx.int64), + attention_mask=tlx.convert_to_tensor(attention_mask, dtype=tlx.int64), + ) + labels = [graph.y for graph in graphs] + zero_predictions = [[0.0] * spec.output_dim for _ in graphs] + metrics = compute_task_metrics(spec, labels, zero_predictions) + return { + "dataset": spec.canonical_name, + "model": args.model, + "num_graphs": len(graphs), + "input_shape": tlx.get_tensor_shape(outputs["last_hidden_state"])[:2], + "logits_shape": tlx.get_tensor_shape(outputs["logits"]), + "mlm_logits_shape": tlx.get_tensor_shape(outputs["mlm_logits"]), + "num_mlm_labels": sum(1 for row in mlm_labels for value in row if value != -100), + "metrics": metrics, + } + + +def build_arg_parser() -> argparse.ArgumentParser: + parser = argparse.ArgumentParser(description="GraphTokenizer training and smoke-test entrypoint.") + parser.add_argument( + "--dataset", + default="qm9", + help="One of: QM9, OGBG-molhiv/molhiv, Peptides-func/p-func, Peptides-struct/p-struct.", + ) + parser.add_argument("--datasets", default=None, help="Comma-separated dataset matrix, e.g. qm9,molhiv,p-func,p-struct.") + parser.add_argument("--model", default=None, choices=SUPPORTED_MODELS) + parser.add_argument("--models", default=None, help="Comma-separated model matrix, e.g. bert,gte.") + parser.add_argument("--protocol", default="gammagl", choices=("gammagl", "paper")) + parser.add_argument( + "--paper-cache-root", + default=None, + help="Tokenizer/serialization cache root.", + ) + parser.add_argument( + "--paper-amp", default=None, choices=("off", "fp16", "bf16"), + help="Optional CUDA autocast mode; strict paper default is off.", + ) + parser.add_argument( + "--paper-tf32", action=argparse.BooleanOptionalAction, default=None, + help="Enable TF32 matmul explicitly; strict paper default is disabled.", + ) + parser.add_argument( + "--paper-loss", default=None, + choices=("default", "l1", "huber", "focal"), + help="Optional supervised loss ablation; default preserves the paper loss.", + ) + parser.add_argument( + "--molhiv-pos-weight", action=argparse.BooleanOptionalAction, + default=None, + help="Weight molhiv positives from training labels only.", + ) + parser.add_argument( + "--target-property", + default=None, + help="Legacy option rejected by strict paper mode, which jointly predicts all 16 QM9 targets.", + ) + parser.add_argument("--device", default="cuda", help="Torch device used by --protocol paper.") + parser.add_argument( + "--allow-random-gte-init", action="store_true", + help=("Use randomly initialized GraphGTE. This disables paper " + "reproduction status and is never the default."), + ) + parser.add_argument( + "--gte-cache-dir", default=None, + help="Optional cache for the pinned official GTE checkpoint.", + ) + parser.add_argument("--num-workers", type=int, default=1) + parser.add_argument("--data-root", default="data") + parser.add_argument("--output-dir", default="logs/graph_tokenizer") + parser.add_argument("--runs", type=int, default=None) + parser.add_argument("--pretrain-epoch", type=int, default=1) + parser.add_argument("--pretrain-lr", type=float, default=1e-4) + parser.add_argument("--pretrain-warmup-ratio", type=float, default=0.12) + parser.add_argument("--pretrain-max-grad-norm", type=float, default=2.0) + parser.add_argument("--n-epoch", type=int, default=1) + parser.add_argument("--finetune-warmup-ratio", type=float, default=0.025) + parser.add_argument("--finetune-max-grad-norm", type=float, default=0.5) + parser.add_argument("--patience", type=int, default=20) + parser.add_argument("--batch-size", type=int, default=32) + parser.add_argument("--lr", type=float, default=1e-4) + parser.add_argument("--weight-decay", type=float, default=0.0) + parser.add_argument("--mask-prob", type=float, default=0.15) + parser.add_argument("--pooling", default="mean", choices=("mean", "cls")) + parser.add_argument("--num-merges", type=int, default=2000) + parser.add_argument("--min-frequency", type=int, default=2) + parser.add_argument("--bpe-backend", default="auto", choices=("python", "auto", "cpp")) + parser.add_argument("--serialization", default="feuler", choices=("feuler", "eulerian")) + parser.add_argument("--max-length", type=int, default=128) + parser.add_argument("--vocab-size", type=int, default=4096) + parser.add_argument("--hidden-size", type=int, default=None) + parser.add_argument("--num-hidden-layers", type=int, default=None) + parser.add_argument("--num-attention-heads", type=int, default=None) + parser.add_argument("--intermediate-size", type=int, default=None) + parser.add_argument("--max-position-embeddings", type=int, default=8192) + parser.add_argument("--seed", type=int, default=None) + parser.add_argument("--no-save-checkpoint", dest="save_checkpoint", action="store_false") + parser.add_argument("--resume", action="store_true", help="Skip dataset/model matrix entries whose summary JSON already exists.") + parser.add_argument("--keep-going", action="store_true", help="Record failed matrix entries and continue with the remaining entries.") + parser.add_argument("--preflight", action="store_true", help="Validate datasets, model names, and BPE backend without training.") + parser.add_argument( + "--prepare-data", action="store_true", + help="Single-process download, verify, materialize, and split-check paper datasets.", + ) + parser.add_argument("--smoke", action="store_true", help="Run a tiny synthetic end-to-end pipeline.") + parser.add_argument("--list-datasets", action="store_true") + parser.set_defaults(save_checkpoint=True) + return parser + + +def main(argv=None): + args = parse_args(argv) + if args.list_datasets: + for spec in DATASET_SPECS: + print(f"{spec.paper_name}: aliases={','.join(spec.aliases)} metric={spec.metric}") + return None + if args.preflight: + report = preflight_check(args) + print(json.dumps(report, indent=2)) + if report["status"] != "ok": + raise SystemExit(1) + return report + if args.prepare_data: + datasets = parse_csv_values(args.datasets) or [args.dataset] + report = {"status": "ok", "datasets": prepare_paper_datasets(args.data_root, datasets)} + print(json.dumps(report, indent=2)) + return report + if args.smoke: + if args.protocol == "paper": + raise ValueError("Paper protocol does not support synthetic --smoke data.") + summary = run_smoke(args) + print(summary) + return summary + if args.protocol == "paper": + if args.datasets or args.models: + raise ValueError( + "Paper protocol runs one dataset and encoder at a time; " + "use --dataset and optional --model.") + summary = run_paper_experiments(args) + print(summary) + return summary + if args.datasets or args.models: + summary = run_experiment_matrix(args) + else: + summary = run_experiments(args) + print(summary) + return summary + + +if __name__ == "__main__": + main() diff --git a/examples/graph_tokenizer/paper_protocol.py b/examples/graph_tokenizer/paper_protocol.py new file mode 100644 index 000000000..b8c33ece2 --- /dev/null +++ b/examples/graph_tokenizer/paper_protocol.py @@ -0,0 +1,1573 @@ +"""Torch-only training protocol aligned with the GraphTokenizer paper.""" + +from __future__ import annotations + +import contextlib +import hashlib +import json +import math +import os +import pickle +import random +import time +from pathlib import Path +from typing import Any, Dict, List, Sequence, Tuple + + +PAPER_RUNTIME_REQUIREMENTS = { + "torch": "2.1.2", + "cuda": "12.1", + "dgl": "2.4.0", + "torch_geometric": "2.4.0", +} + +PAPER_ARCHITECTURES = { + "bert": { + "hidden_size": 512, + "intermediate_size": 2048, + "num_hidden_layers": 4, + "num_attention_heads": 4, + "max_position_embeddings": 8096, + "parameter_estimate": 16_000_000, + }, + "gte": { + "hidden_size": 768, + "intermediate_size": 3072, + "num_hidden_layers": 12, + "num_attention_heads": 12, + "parameter_estimate": 115_000_000, + }, +} + + +def _estimate_training_peak_bytes( + encoder_type: str, micro_batch_size: int, + sequence_length: int, activation_bytes_per_element: int = 4) -> int: + """Conservative optimizer/parameter and saved-activation estimate.""" + architecture = PAPER_ARCHITECTURES[str(encoder_type).lower()] + batch_size = int(micro_batch_size) + sequence_length = int(sequence_length) + layers = int(architecture["num_hidden_layers"]) + heads = int(architecture["num_attention_heads"]) + hidden_size = int(architecture["hidden_size"]) + intermediate_size = int(architecture["intermediate_size"]) + activation_bytes_per_element = int(activation_bytes_per_element) + if activation_bytes_per_element not in {2, 3, 4}: + raise ValueError( + "activation_bytes_per_element must be 2, 3, or 4.") + # FP32 parameters, gradients, and two AdamW moments. + parameter_bytes = int(architecture["parameter_estimate"]) * 16 + # Autograd retains attention probabilities and related score tensors. + attention_bytes = ( + 2 * batch_size * heads * sequence_length * sequence_length + * layers * activation_bytes_per_element + ) + # QKV/residual/normalization/FFN intermediates, including gated GTE FFN. + activation_width = 8 * hidden_size + 3 * intermediate_size + activation_bytes = ( + batch_size * sequence_length * layers * activation_width + * activation_bytes_per_element + ) + return int(parameter_bytes + attention_bytes + activation_bytes) + + +def _resolve_memory_plan( + encoder_type: str, + effective_batch_size: int, + sequence_length: int, + available_bytes: int, + safety_fraction: float = 0.8, + activation_bytes_per_element: int = 4) -> Dict[str, Any]: + """Choose a divisor micro-batch while preserving the paper batch size.""" + encoder_type = str(encoder_type).lower() + if encoder_type not in PAPER_ARCHITECTURES: + raise ValueError("encoder_type must be 'bert' or 'gte'.") + effective_batch_size = int(effective_batch_size) + sequence_length = int(sequence_length) + available_bytes = int(available_bytes) + if effective_batch_size <= 0 or sequence_length <= 0: + raise ValueError("Batch size and sequence length must be positive.") + if available_bytes <= 0: + raise ValueError("available_bytes must be positive.") + if not 0 < float(safety_fraction) <= 1: + raise ValueError("safety_fraction must be in (0, 1].") + budget_bytes = int(available_bytes * float(safety_fraction)) + candidates = [ + value for value in range(effective_batch_size, 0, -1) + if effective_batch_size % value == 0 + ] + for micro_batch_size in candidates: + estimated = _estimate_training_peak_bytes( + encoder_type, micro_batch_size, sequence_length, + activation_bytes_per_element=activation_bytes_per_element) + if estimated <= budget_bytes: + return { + "effective_batch_size": effective_batch_size, + "micro_batch_size": micro_batch_size, + "gradient_accumulation_steps": ( + effective_batch_size // micro_batch_size), + "sequence_length": sequence_length, + "available_bytes": available_bytes, + "budget_bytes": budget_bytes, + "estimated_peak_bytes": estimated, + "safety_fraction": float(safety_fraction), + "activation_bytes_per_element": int( + activation_bytes_per_element), + } + minimum = _estimate_training_peak_bytes( + encoder_type, 1, sequence_length, + activation_bytes_per_element=activation_bytes_per_element) + raise RuntimeError( + "Estimated dense attention training memory is unsafe even with " + f"micro_batch_size=1: estimated={minimum / 1024 ** 3:.2f} GiB, " + f"budget={budget_bytes / 1024 ** 3:.2f} GiB, " + f"encoder={encoder_type}, sequence_length={sequence_length}.") + + +def _load_paper_model_classes(): + from gammagl.models.graph_bert import GraphBERT + from gammagl.models.graph_gte import GraphGTE + + return GraphBERT, GraphGTE + + +def require_paper_runtime(require_cuda: bool = True): + if os.environ.get("TL_BACKEND", "torch").lower() != "torch": + raise RuntimeError("Paper protocol requires TL_BACKEND=torch.") + import tensorlayerx as tlx + if str(tlx.BACKEND).lower() != "torch": + raise RuntimeError( + "Paper protocol requires the TensorLayerX PyTorch backend; " + f"received {tlx.BACKEND!r}.") + try: + import torch + except ImportError as error: + raise RuntimeError("Paper protocol requires PyTorch.") from error + if require_cuda and not torch.cuda.is_available(): + raise RuntimeError("Paper protocol requires a CUDA-enabled PyTorch runtime.") + return torch + + +def _create_paper_model( + encoder_type: str, + vocab_size: int, + pad_token_id: int, + task_type: str, + output_dim: int, + pooling: str = "mean", + model_config: Dict[str, Any] | None = None, + strict_architecture: bool = True, + allow_random_gte_init: bool = False, + pretrained_cache_dir: str | None = None): + GraphBERT, GraphGTE = _load_paper_model_classes() + encoder_type = str(encoder_type).lower() + if encoder_type not in {"bert", "gte"}: + raise ValueError("encoder_type must be 'bert' or 'gte'.") + model_class = GraphBERT if encoder_type == "bert" else GraphGTE + kwargs = { + "vocab_size": int(vocab_size), + "output_dim": int(output_dim), + "pad_token_id": int(pad_token_id), + "task_type": task_type, + "pooling": pooling, + } + if not strict_architecture: + allowed = { + "hidden_size", + "num_hidden_layers", + "num_attention_heads", + "intermediate_size", + "max_position_embeddings", + "dropout_rate", + "attention_dropout_rate", + "layer_norm_eps", + "task_dropout", + "rope_theta", + "rope_scaling_factor", + } + kwargs.update({ + key: value + for key, value in dict(model_config or {}).items() + if key in allowed + }) + if encoder_type == "gte" and not allow_random_gte_init: + # A paper GTE run is a reproduction only with the pinned official + # encoder. Any download/config/conversion error intentionally escapes. + return GraphGTE.from_pretrained(cache_dir=pretrained_cache_dir, **kwargs) + if encoder_type == "bert" and strict_architecture: + kwargs["max_position_embeddings"] = PAPER_ARCHITECTURES["bert"][ + "max_position_embeddings"] + model = model_class(**kwargs) + if encoder_type == "gte": + model.pretrained_manifest = { + "pretrained": False, + "reproduction": False, + "graph_embedding_init": "new_truncated_normal(stddev=0.02)", + "task_head_init": "new_truncated_normal(stddev=0.02)", + } + return model + + +def _require_fp32_parameters(torch, model) -> None: + non_fp32 = [ + name + for name, parameter in model.named_parameters() + if parameter.is_floating_point() and parameter.dtype != torch.float32 + ] + if non_fp32: + preview = ", ".join(non_fp32[:3]) + raise RuntimeError( + "Paper protocol requires FP32 model parameters before AdamW; " + f"received non-FP32 parameters: {preview}.") + + +def _torch(): + return require_paper_runtime(require_cuda=True) + + +def validate_paper_runtime_versions(versions: Dict[str, Any]) -> None: + for field, expected in PAPER_RUNTIME_REQUIREMENTS.items(): + actual = str(versions.get(field, "")) + comparable = actual.split("+", 1)[0] if field in {"torch", "dgl"} else actual + if comparable != expected: + raise RuntimeError( + f"Paper runtime requires {field}=={expected}; received {actual}.") + + +def paper_runtime_status() -> Dict[str, Any]: + torch = _torch() + try: + import dgl + import numpy as np + import ogb + import torch_geometric + except ImportError as error: + raise RuntimeError( + "Paper protocol requires DGL, PyTorch Geometric, numpy and OGB.") from error + device_index = torch.cuda.current_device() + status = { + "status": "ok", + "torch": torch.__version__, + "cuda": torch.version.cuda, + "dgl": dgl.__version__, + "torch_geometric": torch_geometric.__version__, + "cuda_device": torch.cuda.get_device_name(device_index), + "numpy": np.__version__, + "ogb": ogb.__version__, + } + validate_paper_runtime_versions(status) + status["requirements"] = dict(PAPER_RUNTIME_REQUIREMENTS) + return status + + +def validate_paper_args(args, spec) -> None: + if spec.canonical_name not in {"qm9", "molhiv", "peptides-struct"}: + raise ValueError( + "Paper protocol supports only QM9, OGBG-molhiv and Peptides-struct.") + if spec.canonical_name == "qm9": + target_property = getattr(args, "target_property", None) + if target_property is not None: + raise ValueError( + "QM9 strict paper protocol is a joint 16-target regression; " + "do not pass --target-property.") + requested_model = getattr(args, "model", None) + if requested_model and requested_model not in {"bert", "gte"}: + raise ValueError("Paper protocol model must be 'bert' or 'gte'.") + if requested_model: + expected_position_limit = PAPER_ARCHITECTURES.get( + requested_model, {}).get("max_position_embeddings") + if expected_position_limit is not None and int(getattr( + args, "max_position_embeddings", expected_position_limit)) != expected_position_limit: + raise ValueError( + "Paper protocol requires " + f"--max-position-embeddings={expected_position_limit} for " + f"{requested_model}.") + + +def paper_run_seeds(args) -> List[int]: + runs = getattr(args, "runs", None) + runs = 5 if runs is None else int(runs) + if runs != 5: + raise ValueError("Paper protocol requires exactly five independent runs.") + start_seed = 42 if getattr(args, "seed", None) is None else int(args.seed) + return [start_seed + index for index in range(runs)] + + +def select_best_epoch(history: Sequence[Dict[str, Any]], higher_is_better: bool) -> int: + if not history: + raise ValueError("Cannot select a checkpoint from an empty training history.") + key = lambda entry: float(entry["val_metric"]) + return int((max if higher_is_better else min)(history, key=key)["epoch"]) + + +def _should_skip_finetuning( + resume_phase, stale_epochs: int, patience: int) -> bool: + return resume_phase == "finetune_complete" or ( + resume_phase == "finetune" + and int(stale_epochs) >= int(patience) + ) + + +def paper_training_options(args) -> Dict[str, Any]: + options = { + "encoder": args.model, + "serialization": args.serialization, + "pretrain_epochs": args.pretrain_epoch, + "finetune_epochs": args.n_epoch, + "pretrain_lr": args.pretrain_lr, + "finetune_lr": args.lr, + "batch_size": args.batch_size, + "weight_decay": args.weight_decay, + "pretrain_warmup_ratio": args.pretrain_warmup_ratio, + "finetune_warmup_ratio": args.finetune_warmup_ratio, + "pretrain_max_grad_norm": args.pretrain_max_grad_norm, + "finetune_max_grad_norm": args.finetune_max_grad_norm, + "mask_prob": args.mask_prob, + "patience": args.patience, + "max_position_embeddings": args.max_position_embeddings, + } + cli_overrides = { + "amp_dtype": getattr(args, "paper_amp", None), + "allow_tf32": getattr(args, "paper_tf32", None), + "training_loss": getattr(args, "paper_loss", None), + "molhiv_pos_weight": getattr(args, "molhiv_pos_weight", None), + } + options.update({ + key: value for key, value in cli_overrides.items() + if value is not None + }) + return options + + +def set_paper_seed(seed: int, torch) -> None: + random.seed(seed) + try: + import numpy as np + + np.random.seed(seed) + except ImportError: + pass + torch.manual_seed(seed) + torch.cuda.manual_seed_all(seed) + torch.backends.cudnn.deterministic = True + torch.backends.cudnn.benchmark = False + + +def _pad_sequences(sequences: Sequence[Sequence[int]], max_length: int, pad_token_id: int): + if max_length <= 0: + raise ValueError("max_length must be positive.") + if any(len(sequence) > max_length for sequence in sequences): + raise ValueError( + f"A token sequence exceeds the configured max_length={max_length}.") + input_ids = [list(map(int, sequence)) + [pad_token_id] * (max_length - len(sequence)) + for sequence in sequences] + masks = [[1] * len(sequence) + [0] * (max_length - len(sequence)) + for sequence in sequences] + return input_ids, masks + + +def _encode_splits( + tokenizer, splits, max_position_embeddings: int, + encoding_batch_size: int = 2048): + encoded = {} + global_max_length = 1 + for split_name, graphs in splits.items(): + if not graphs: + raise ValueError( + f"Paper protocol requires a non-empty {split_name} split.") + if hasattr(tokenizer, "batch_encode_graphs"): + sequences = [] + encoding_batch_size = max(1, int(encoding_batch_size)) + for start in range(0, len(graphs), encoding_batch_size): + graph_batch = graphs[start:start + encoding_batch_size] + results = tokenizer.batch_encode_graphs(graph_batch) + if len(results) != len(graph_batch): + raise RuntimeError( + "Tokenizer batch encoding returned the wrong count.") + sequences.extend(result.input_ids for result in results) + else: + sequences = [ + tokenizer.encode_graph(graph).input_ids for graph in graphs + ] + maximum = max((len(sequence) for sequence in sequences), default=1) + if maximum > max_position_embeddings: + raise ValueError( + f"{split_name} sequence length {maximum} exceeds " + f"max_position_embeddings={max_position_embeddings}.") + global_max_length = max(global_max_length, maximum) + encoded[split_name] = { + "input_ids": [list(map(int, sequence)) for sequence in sequences], + "labels": [list(graph.y) for graph in graphs], + } + return encoded, global_max_length + + +def _load_or_encode_splits( + tokenizer, splits, max_position_embeddings: int, + cache_path, cache_key: str): + cache_path = Path(cache_path) + if cache_path.is_file(): + try: + with cache_path.open("rb") as handle: + cached = pickle.load(handle) + if ( + cached.get("cache_version") == 1 + and cached.get("cache_key") == str(cache_key) + and cached.get("max_position_embeddings") + == int(max_position_embeddings)): + return cached["encoded"], int(cached["global_max_length"]) + except (OSError, EOFError, AttributeError, KeyError, pickle.UnpicklingError): + pass + encoded, global_max_length = _encode_splits( + tokenizer, splits, int(max_position_embeddings)) + cache_path.parent.mkdir(parents=True, exist_ok=True) + temporary = cache_path.with_name( + f".{cache_path.name}.tmp-{os.getpid()}") + with temporary.open("wb") as handle: + pickle.dump({ + "cache_version": 1, + "cache_key": str(cache_key), + "max_position_embeddings": int(max_position_embeddings), + "global_max_length": int(global_max_length), + "encoded": encoded, + }, handle, protocol=pickle.HIGHEST_PROTOCOL) + os.replace(temporary, cache_path) + return encoded, global_max_length + + +def _encoded_splits_cache_key(tokenizer, splits, max_position_embeddings: int) -> str: + split_digests = {} + for split_name, graphs in splits.items(): + digest = hashlib.sha256() + for graph in graphs: + fields = [] + for field_name in ( + "edge_index", "x", "edge_attr", "y", "num_nodes", + "tokens"): + value = getattr(graph, field_name, None) + if hasattr(value, "detach"): + value = value.detach() + if hasattr(value, "cpu"): + value = value.cpu() + if hasattr(value, "tolist"): + value = value.tolist() + fields.append((field_name, value)) + digest.update(pickle.dumps(fields, protocol=4)) + split_digests[split_name] = digest.hexdigest() + state = { + "cache_version": 1, + "tokenizer_cache_key": getattr(tokenizer, "_cache_key", None), + "serializer": getattr(tokenizer.serializer, "name", None), + "frequency_map": sorted(tokenizer.serializer.frequency_map.items()), + "merge_rules": list(tokenizer.bpe.codebook.merge_rules), + "vocab_size": int(tokenizer.bpe.codebook.vocab_size), + "split_sizes": { + split_name: len(graphs) for split_name, graphs in splits.items() + }, + "split_digests": split_digests, + "max_position_embeddings": int(max_position_embeddings), + } + return hashlib.sha256(pickle.dumps(state, protocol=4)).hexdigest() + + +def _prepare_labels(encoded, spec, target_property): + if spec.canonical_name == "qm9": + output_dim = int(spec.output_dim) + target_properties = list(spec.label_keys) + if len(target_properties) != output_dim: + raise ValueError( + f"QM9 requires exactly {output_dim} target property names; " + f"received {len(target_properties)}.") + if not encoded["train"]["labels"]: + raise ValueError("QM9 training split cannot be empty.") + for split_name, split in encoded.items(): + for row_index, row in enumerate(split["labels"]): + if len(row) != output_dim: + raise ValueError( + f"QM9 {split_name} label row {row_index} must contain " + f"exactly {output_dim} targets; received {len(row)}.") + if not all(math.isfinite(float(value)) for value in row): + raise ValueError( + f"QM9 {split_name} label row {row_index} contains a " + "non-finite target.") + train_rows = [list(map(float, row)) for row in encoded["train"]["labels"]] + means = [ + sum(row[index] for row in train_rows) / len(train_rows) + for index in range(output_dim) + ] + variances = [ + sum((row[index] - means[index]) ** 2 for row in train_rows) + / len(train_rows) + for index in range(output_dim) + ] + stds = [max(math.sqrt(variance), 1e-12) for variance in variances] + for split in encoded.values(): + split["labels"] = [ + [ + (float(row[index]) - means[index]) / stds[index] + for index in range(output_dim) + ] + for row in split["labels"] + ] + return { + "mean": means, + "std": stds, + "target_properties": target_properties, + }, output_dim + return None, spec.output_dim + + +def _denormalize_labels(torch, values, normalizer): + means = torch.as_tensor( + normalizer["mean"], dtype=values.dtype, device=values.device) + stds = torch.as_tensor( + normalizer["std"], dtype=values.dtype, device=values.device) + if values.ndim < 1 or values.shape[-1] != len(means): + raise ValueError( + "Normalized label width does not match the QM9 normalizer.") + return values * stds + means + + +def _metric_semantics(spec, normalizer): + """Describe the training-loss and reported-metric spaces explicitly.""" + if spec.canonical_name == "qm9": + normalization = {"type": "train_split_zscore"} + if normalizer is not None: + normalization.update({ + "mean": list(normalizer["mean"]), + "std": list(normalizer["std"]), + }) + return { + "loss_space": "standardized", + "metric_space": "raw", + "target_normalization": normalization, + } + if spec.canonical_name == "molhiv": + return { + "loss_space": "logits", + "metric_space": "probability", + "target_normalization": {"type": "none"}, + } + if spec.canonical_name == "peptides-struct": + return { + "loss_space": "raw", + "metric_space": "raw", + "target_normalization": {"type": "none"}, + } + return { + "loss_space": "logits", + "metric_space": "probability", + "target_normalization": {"type": "none"}, + } + + +class _LengthBucketBatchSampler: + """Shuffle examples, then form locally length-sorted dynamic batches.""" + + def __init__( + self, torch, lengths, batch_size: int, shuffle: bool, + pool_size_multiplier: int = 50): + self.torch = torch + self.lengths = [int(length) for length in lengths] + self.batch_size = int(batch_size) + self.shuffle = bool(shuffle) + self.pool_size = max( + self.batch_size, self.batch_size * int(pool_size_multiplier)) + + def __len__(self): + return (len(self.lengths) + self.batch_size - 1) // self.batch_size + + def __iter__(self): + if self.shuffle: + indices = self.torch.randperm(len(self.lengths)).tolist() + else: + indices = list(range(len(self.lengths))) + batches = [] + for start in range(0, len(indices), self.pool_size): + pool = indices[start:start + self.pool_size] + if self.shuffle: + pool.sort(key=lambda index: self.lengths[index], reverse=True) + for offset in range(0, len(pool), self.batch_size): + batches.append(pool[offset:offset + self.batch_size]) + if self.shuffle and len(batches) > 1: + order = self.torch.randperm(len(batches)).tolist() + batches = [batches[index] for index in order] + yield from batches + + +def _make_loader( + torch, + encoded_split, + batch_size: int, + shuffle: bool, + pad_token_id: int = 0, + num_workers: int = 0, + pin_memory: bool = False, + bucket_by_length: bool = True): + input_ids = encoded_split["input_ids"] + labels = encoded_split["labels"] + if len(input_ids) != len(labels): + raise ValueError("Encoded input and label counts must match.") + if not input_ids: + raise ValueError("Cannot create a paper DataLoader from an empty split.") + dataset = list(zip(input_ids, labels)) + + def collate_batch(batch): + maximum = max(len(sequence) for sequence, _ in batch) + padded = torch.full( + (len(batch), maximum), int(pad_token_id), dtype=torch.long) + attention_mask = torch.zeros_like(padded) + for row, (sequence, _) in enumerate(batch): + length = len(sequence) + if length: + padded[row, :length] = torch.as_tensor( + sequence, dtype=torch.long) + attention_mask[row, :length] = 1 + batch_labels = torch.as_tensor( + [label for _, label in batch], dtype=torch.float32) + return padded, attention_mask, batch_labels + + workers = max(0, int(num_workers)) + worker_seed_generator = torch.Generator() + worker_seed_generator.manual_seed(0) + loader_kwargs = { + "num_workers": workers, + "pin_memory": bool(pin_memory), + "collate_fn": collate_batch, + "generator": worker_seed_generator, + } + if workers: + loader_kwargs.update({ + "persistent_workers": True, + "prefetch_factor": 2, + }) + if bucket_by_length: + loader_kwargs["batch_sampler"] = _LengthBucketBatchSampler( + torch, + [len(sequence) for sequence in input_ids], + batch_size=int(batch_size), + shuffle=shuffle, + ) + else: + loader_kwargs.update({ + "batch_size": int(batch_size), + "shuffle": bool(shuffle), + }) + loader = torch.utils.data.DataLoader(dataset, **loader_kwargs) + loader.dynamic_padding = True + loader.worker_seed_generator = worker_seed_generator + return loader + + +def _linear_warmup_scheduler(torch, optimizer, total_steps: int, warmup_ratio: float): + warmup_steps = max(1, int(total_steps * float(warmup_ratio))) + + def multiplier(step): + if step < warmup_steps: + return float(step + 1) / warmup_steps + return max(0.0, float(total_steps - step) / max(1, total_steps - warmup_steps)) + + return torch.optim.lr_scheduler.LambdaLR(optimizer, multiplier) + + +def _mask_for_mlm(torch, input_ids, attention_mask, tokenizer, mask_prob: float): + labels = input_ids.clone() + maskable = attention_mask.bool() + for token_id in ( + tokenizer.special_tokens.pad_token_id, + tokenizer.special_tokens.cls_token_id, + tokenizer.special_tokens.sep_token_id, + tokenizer.special_tokens.component_sep_token_id, + ): + maskable &= input_ids.ne(int(token_id)) + selected = torch.rand_like(input_ids, dtype=torch.float32).lt(float(mask_prob)) & maskable + if float(mask_prob) > 0: + flat_maskable = maskable.reshape(-1) + first_maskable = flat_maskable & flat_maskable.cumsum(dim=0).eq(1) + selected |= first_maskable.reshape_as(selected) & ~selected.any() + labels[~selected] = -100 + masked = input_ids.clone() + masked[selected] = int(tokenizer.special_tokens.mask_token_id) + return masked, labels + + +def _loss( + torch, logits, labels, spec, loss_name: str = "default", + pos_weight=None, focal_gamma: float = 2.0): + logits = logits.float() + labels = labels.float() + loss_name = str(loss_name).lower() + if spec.canonical_name == "molhiv": + if loss_name not in {"default", "focal"}: + raise ValueError("OGBG-molhiv loss must be 'default' or 'focal'.") + logits = logits.view(-1) + labels = labels.view(-1) + bce = torch.nn.functional.binary_cross_entropy_with_logits( + logits, labels, pos_weight=pos_weight, reduction="none") + if loss_name == "focal": + probabilities = torch.sigmoid(logits) + probability_of_target = ( + probabilities * labels + (1 - probabilities) * (1 - labels)) + bce = (1 - probability_of_target).pow(float(focal_gamma)) * bce + return bce.mean() + valid = None + if spec.canonical_name == "peptides-struct": + valid = ~torch.isnan(labels) + if not torch.any(valid): + raise ValueError("Peptides-struct batch contains no valid labels.") + logits = logits[valid] + labels = labels[valid] + if loss_name == "default": + return torch.nn.functional.mse_loss(logits, labels) + if loss_name == "l1": + return torch.nn.functional.l1_loss(logits, labels) + if loss_name == "huber": + return torch.nn.functional.smooth_l1_loss(logits, labels) + raise ValueError("Regression loss must be 'default', 'l1', or 'huber'.") + + +def _resolve_precision_options(torch, device, preset) -> Dict[str, Any]: + amp_dtype = str(preset.get("amp_dtype", "off")).lower() + if amp_dtype not in {"off", "fp16", "bf16"}: + raise ValueError("amp_dtype must be 'off', 'fp16', or 'bf16'.") + if amp_dtype != "off" and device.type != "cuda": + raise ValueError("Mixed precision requires a CUDA device.") + torch_dtype = { + "off": None, + "fp16": torch.float16, + "bf16": torch.bfloat16, + }[amp_dtype] + return { + "amp_dtype": amp_dtype, + "torch_dtype": torch_dtype, + "grad_scaler": None, + "allow_tf32": bool(preset.get("allow_tf32", False)), + } + + +def _cuda_memory_query_device(torch, device): + """Return an explicit CUDA index for APIs that reject bare ``cuda``.""" + if device.type == "cuda" and device.index is None: + return torch.cuda.current_device() + return device + + +def _new_grad_scaler(torch, amp_dtype: str): + if str(amp_dtype).lower() != "fp16": + return None + return torch.cuda.amp.GradScaler(enabled=True) + + +def _autocast_context(torch, device, torch_dtype): + if torch_dtype is None: + return contextlib.nullcontext() + return torch.autocast( + device_type=device.type, dtype=torch_dtype, enabled=True) + + +def _set_paper_model_mode(model, training: bool) -> None: + """Use TensorLayerX modes, with a PyTorch fallback for test doubles.""" + tlx_method = getattr( + model, "set_train" if training else "set_eval", None) + if callable(tlx_method): + tlx_method() + elif training: + model.train() + else: + model.eval() + + +def _epoch_runtime_metrics( + torch, device, num_examples: int, seconds: float) -> Dict[str, float]: + seconds = float(seconds) + metrics = { + "seconds": seconds, + "examples_per_second": float(num_examples) / max(seconds, 1e-12), + } + if device.type == "cuda": + metrics["peak_cuda_memory_gib"] = ( + float(torch.cuda.max_memory_allocated(device)) / 1024 ** 3) + return metrics + + +def _require_finite_tensor(torch, value, name: str) -> None: + if not torch.isfinite(value).all(): + raise FloatingPointError(f"Non-finite {name} detected.") + + +def _require_finite_scalar(value, name: str) -> float: + value = float(value) + if not math.isfinite(value): + raise FloatingPointError(f"Non-finite {name} detected.") + return value + + +def _average_mae(y_true, y_pred): + import numpy as np + + y_true = np.asarray(y_true, dtype=float) + y_pred = np.asarray(y_pred, dtype=float) + task_scores = [] + for task_index in range(y_true.shape[1]): + valid = ~np.isnan(y_true[:, task_index]) + if np.any(valid): + task_scores.append(float(np.abs( + y_true[valid, task_index] - y_pred[valid, task_index]).mean())) + return sum(task_scores) / len(task_scores) if task_scores else float("nan") + + +def compute_paper_metric(spec, y_true, y_pred) -> float: + import numpy as np + + y_true = np.asarray(y_true, dtype=float) + y_pred = np.asarray(y_pred, dtype=float) + if spec.canonical_name == "molhiv": + try: + from ogb.graphproppred import Evaluator + except ImportError as error: + raise RuntimeError("Paper OGBG-molhiv evaluation requires the ogb package.") from error + evaluator = Evaluator(name="ogbg-molhiv") + return float(evaluator.eval({ + "y_true": np.asarray(y_true, dtype=float), + "y_pred": np.asarray(y_pred, dtype=float), + })["rocauc"]) + if spec.canonical_name == "peptides-struct": + return _average_mae(y_true, y_pred) + return float(np.abs(y_true - y_pred).mean()) + + +def compute_paper_metric_details(spec, y_true, y_pred) -> Dict[str, Any]: + import numpy as np + + y_true = np.asarray(y_true, dtype=float) + y_pred = np.asarray(y_pred, dtype=float) + details = {"metric": compute_paper_metric(spec, y_true, y_pred)} + if spec.canonical_name not in {"qm9", "peptides-struct"}: + return details + if y_true.ndim == 1: + y_true = y_true.reshape(-1, 1) + y_pred = y_pred.reshape(-1, 1) + configured_names = list(getattr(spec, "label_keys", ())) + if len(configured_names) == y_true.shape[1]: + target_names = configured_names + else: + target_names = [f"target_{index}" for index in range(y_true.shape[1])] + per_target = {} + for index, target_name in enumerate(target_names): + valid = np.isfinite(y_true[:, index]) & np.isfinite(y_pred[:, index]) + per_target[target_name] = ( + float(np.abs( + y_true[valid, index] - y_pred[valid, index]).mean()) + if valid.any() else float("nan")) + details["per_target_mae"] = per_target + return details + + +def _evaluate( + torch, model, loader, spec, device, normalizer, torch_dtype=None): + _set_paper_model_mode(model, training=False) + total_loss = 0.0 + total_examples = 0 + targets, predictions = [], [] + non_blocking = bool(getattr(loader, "pin_memory", False)) + with torch.no_grad(): + for input_ids, attention_mask, labels in loader: + input_ids = input_ids.to(device, non_blocking=non_blocking) + attention_mask = attention_mask.to( + device, non_blocking=non_blocking) + labels = labels.to(device, non_blocking=non_blocking) + with _autocast_context(torch, device, torch_dtype): + logits = model(input_ids, attention_mask, task="supervised") + _require_finite_tensor(torch, logits, "supervised logits") + loss = _loss(torch, logits, labels, spec) + _require_finite_tensor(torch, loss, "evaluation loss") + total_loss += loss.item() * len(input_ids) + total_examples += len(input_ids) + metric_logits = logits.float() + if spec.canonical_name == "molhiv": + predictions.append(torch.sigmoid(metric_logits).cpu()) + else: + predictions.append(metric_logits.cpu()) + targets.append(labels.cpu()) + y_true = torch.cat(targets, dim=0) + y_pred = torch.cat(predictions, dim=0) + average_loss = _require_finite_scalar( + total_loss / max(1, total_examples), "evaluation loss") + metric_details = compute_paper_metric_details( + spec, y_true.numpy(), y_pred.numpy()) + metric_details["metric_space"] = _metric_semantics( + spec, normalizer)["metric_space"] + if normalizer is not None and spec.canonical_name == "qm9": + raw_y_true = _denormalize_labels(torch, y_true, normalizer) + raw_y_pred = _denormalize_labels(torch, y_pred, normalizer) + raw_details = compute_paper_metric_details( + spec, raw_y_true.numpy(), raw_y_pred.numpy()) + metric_details["metric"] = raw_details["metric"] + metric_details["per_target_mae_raw"] = raw_details["per_target_mae"] + metric = _require_finite_scalar( + metric_details["metric"], + f"{spec.canonical_name} validation metric", + ) + return { + "loss": average_loss, + "metric": metric, + **{ + key: value for key, value in metric_details.items() + if key != "metric" + }, + } + + +def _rng_state(torch): + state = {"python": random.getstate(), "torch": torch.get_rng_state()} + try: + import numpy as np + + state["numpy"] = np.random.get_state() + except ImportError: + pass + if torch.cuda.is_available(): + state["cuda"] = torch.cuda.get_rng_state_all() + return state + + +def _atomic_torch_save(torch, state, path) -> str: + path = Path(path) + path.parent.mkdir(parents=True, exist_ok=True) + temporary = path.with_name(f".{path.name}.tmp-{os.getpid()}") + torch.save(state, temporary) + os.replace(temporary, path) + return str(path) + + +def _torch_load(torch, path, device): + try: + return torch.load(path, map_location=device, weights_only=False) + except TypeError: + return torch.load(path, map_location=device) + + +def save_paper_checkpoint( + path, torch, model, epoch, best_metric, normalizer, + experiment_fingerprint: str | None = None): + """Save the small best-model artifact used only for final evaluation.""" + best_metric = _require_finite_scalar(best_metric, "checkpoint metric") + state = { + "checkpoint_kind": "best_model", + "model": model.state_dict(), + "epoch": int(epoch), + "best_metric": best_metric, + "normalizer": normalizer, + } + if experiment_fingerprint is not None: + state["experiment_fingerprint"] = str(experiment_fingerprint) + return _atomic_torch_save(torch, state, path) + + +def _validate_experiment_fingerprint(state, expected_fingerprint) -> None: + if expected_fingerprint is None: + return + actual = state.get("experiment_fingerprint") + if actual != str(expected_fingerprint): + raise ValueError( + "Checkpoint experiment fingerprint does not match the current " + "data/tokenizer/configuration; use a new output directory or " + "restart without --resume.") + + +def restore_paper_checkpoint( + path, torch, model, device, expected_fingerprint=None): + state = _torch_load(torch, path, device) + if state.get("checkpoint_kind") != "best_model": + raise ValueError(f"Not a best-model checkpoint: {path}") + _validate_experiment_fingerprint(state, expected_fingerprint) + model.load_state_dict(state["model"]) + return state + + +def save_paper_resume_state( + path, torch, model, optimizer, scheduler, phase: str, + epoch: int, extra: Dict[str, Any] | None = None, grad_scaler=None, + experiment_fingerprint: str | None = None): + state = { + "checkpoint_kind": "resume_state", + "phase": str(phase), + "epoch": int(epoch), + "model": model.state_dict(), + "optimizer": optimizer.state_dict(), + "scheduler": scheduler.state_dict(), + "rng_state": _rng_state(torch), + "extra": dict(extra or {}), + } + if grad_scaler is not None: + state["grad_scaler"] = grad_scaler.state_dict() + if experiment_fingerprint is not None: + state["experiment_fingerprint"] = str(experiment_fingerprint) + return _atomic_torch_save(torch, state, path) + + +def _restore_rng_state(torch, state) -> None: + rng_state = state.get("rng_state", {}) + if "python" in rng_state: + random.setstate(rng_state["python"]) + if "torch" in rng_state: + torch.set_rng_state(rng_state["torch"]) + if "cuda" in rng_state and torch.cuda.is_available(): + torch.cuda.set_rng_state_all(rng_state["cuda"]) + if "numpy" in rng_state: + try: + import numpy as np + + np.random.set_state(rng_state["numpy"]) + except ImportError: + pass + + +def restore_paper_resume_state( + path, torch, model, optimizer, scheduler, device, grad_scaler=None, + expected_fingerprint=None): + state = _torch_load(torch, path, device) + if state.get("checkpoint_kind") != "resume_state": + raise ValueError(f"Not a resume-state checkpoint: {path}") + _validate_experiment_fingerprint(state, expected_fingerprint) + model.load_state_dict(state["model"]) + optimizer.load_state_dict(state["optimizer"]) + scheduler.load_state_dict(state["scheduler"]) + _restore_grad_scaler_state(state, grad_scaler) + _restore_rng_state(torch, state) + return state + + +def _restore_grad_scaler_state(state, grad_scaler) -> None: + if grad_scaler is not None and "grad_scaler" in state: + grad_scaler.load_state_dict(state["grad_scaler"]) + + +def _experiment_fingerprint( + preset, spec, encoded_cache_key: str, model_vocab_size: int, + pooling: str, seed: int, run_index: int) -> str: + payload = { + "protocol_version": 2, + "preset": preset, + "dataset": spec.canonical_name, + "task_type": spec.task_type, + "output_dim": int(spec.output_dim), + "encoded_cache_key": str(encoded_cache_key), + "model_vocab_size": int(model_vocab_size), + "pooling": str(pooling), + "seed": int(seed), + "run_index": int(run_index), + } + return hashlib.sha256(json.dumps( + payload, sort_keys=True, separators=(",", ":"), default=str + ).encode("utf-8")).hexdigest() + + +def _optimizer_steps_per_epoch(loader, gradient_accumulation_steps: int) -> int: + accumulation = max(1, int(gradient_accumulation_steps)) + return (len(loader) + accumulation - 1) // accumulation + + +def _molhiv_positive_weight(torch, encoded_train, device): + labels = torch.as_tensor(encoded_train["labels"], dtype=torch.float32) + labels = labels[torch.isfinite(labels)] + positives = labels.eq(1).sum() + negatives = labels.eq(0).sum() + if positives.item() == 0 or negatives.item() == 0: + raise ValueError( + "OGBG-molhiv class weighting requires both classes in train.") + return (negatives.float() / positives.float()).to(device) + + +def _supervised_loss_weight(torch, labels, spec) -> int: + if spec.canonical_name == "peptides-struct": + return int(torch.isfinite(labels).sum().item()) + return int(labels.numel()) + + +def _normalize_accumulated_gradients(model, denominator: int) -> None: + denominator = max(1, int(denominator)) + for parameter in model.parameters(): + if parameter.grad is not None: + parameter.grad.div_(denominator) + + +def _train_supervised( + torch, model, loader, optimizer, scheduler, spec, device, + max_grad_norm, gradient_accumulation_steps: int = 1, + loss_name: str = "default", pos_weight=None, + torch_dtype=None, grad_scaler=None): + _set_paper_model_mode(model, training=True) + non_blocking = bool(getattr(loader, "pin_memory", False)) + accumulation = max(1, int(gradient_accumulation_steps)) + total_batches = len(loader) + total_loss = 0.0 + total_loss_weight = 0 + group_loss_weight = 0 + optimizer.zero_grad(set_to_none=True) + for step, (input_ids, attention_mask, labels) in enumerate(loader): + labels = labels.to(device, non_blocking=non_blocking) + with _autocast_context(torch, device, torch_dtype): + logits = model( + input_ids.to(device, non_blocking=non_blocking), + attention_mask.to(device, non_blocking=non_blocking), + task="supervised", + ) + _require_finite_tensor(torch, logits, "supervised logits") + loss = _loss( + torch, + logits, + labels, + spec, + loss_name=loss_name, + pos_weight=pos_weight, + ) + _require_finite_tensor(torch, loss, "supervised loss") + loss_weight = _supervised_loss_weight(torch, labels, spec) + total_loss += float(loss.detach().item()) * loss_weight + total_loss_weight += loss_weight + group_loss_weight += loss_weight + scaled_loss = loss * loss_weight + if grad_scaler is None: + scaled_loss.backward() + else: + grad_scaler.scale(scaled_loss).backward() + should_step = (step + 1) % accumulation == 0 or step + 1 == total_batches + if not should_step: + continue + if grad_scaler is not None: + grad_scaler.unscale_(optimizer) + _normalize_accumulated_gradients(model, group_loss_weight) + grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), float(max_grad_norm)) + if grad_scaler is None: + _require_finite_tensor(torch, grad_norm, "supervised gradient norm") + optimizer.step() + scheduler.step() + else: + scale_before_step = grad_scaler.get_scale() + grad_scaler.step(optimizer) + grad_scaler.update() + if grad_scaler.get_scale() >= scale_before_step: + scheduler.step() + optimizer.zero_grad(set_to_none=True) + group_loss_weight = 0 + return total_loss / max(1, total_loss_weight) + + +def _train_mlm( + torch, model, loader, optimizer, scheduler, tokenizer, device, + max_grad_norm, mask_prob, gradient_accumulation_steps: int = 1, + torch_dtype=None, grad_scaler=None): + _set_paper_model_mode(model, training=True) + non_blocking = bool(getattr(loader, "pin_memory", False)) + accumulation = max(1, int(gradient_accumulation_steps)) + total_batches = len(loader) + total_loss = 0.0 + total_loss_weight = 0 + group_loss_weight = 0 + optimizer.zero_grad(set_to_none=True) + for step, (input_ids, attention_mask, _) in enumerate(loader): + input_ids = input_ids.to(device, non_blocking=non_blocking) + attention_mask = attention_mask.to( + device, non_blocking=non_blocking) + masked_ids, labels = _mask_for_mlm(torch, input_ids, attention_mask, tokenizer, mask_prob) + with _autocast_context(torch, device, torch_dtype): + logits = model(masked_ids, attention_mask, task="mlm") + _require_finite_tensor(torch, logits, "MLM logits") + loss = torch.nn.functional.cross_entropy( + logits.reshape(-1, logits.shape[-1]), + labels.reshape(-1), + ignore_index=-100, + ) + _require_finite_tensor(torch, loss, "MLM loss") + loss_weight = int(labels.ne(-100).sum().item()) + total_loss += float(loss.detach().item()) * loss_weight + total_loss_weight += loss_weight + group_loss_weight += loss_weight + scaled_loss = loss * loss_weight + if grad_scaler is None: + scaled_loss.backward() + else: + grad_scaler.scale(scaled_loss).backward() + should_step = (step + 1) % accumulation == 0 or step + 1 == total_batches + if not should_step: + continue + if grad_scaler is not None: + grad_scaler.unscale_(optimizer) + _normalize_accumulated_gradients(model, group_loss_weight) + grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), float(max_grad_norm)) + if grad_scaler is None: + _require_finite_tensor(torch, grad_norm, "MLM gradient norm") + optimizer.step() + scheduler.step() + else: + scale_before_step = grad_scaler.get_scale() + grad_scaler.step(optimizer) + grad_scaler.update() + if grad_scaler.get_scale() >= scale_before_step: + scheduler.step() + optimizer.zero_grad(set_to_none=True) + group_loss_weight = 0 + return total_loss / max(1, total_loss_weight) + + +def run_paper_experiment(args, spec, splits, fit_tokenizer): + validate_paper_args(args, spec) + torch = _torch() + preset = paper_training_options(args) + seeds = paper_run_seeds(args) + if not splits["train"]: + raise ValueError("Paper protocol refuses an empty training split.") + args.serialization = preset["serialization"] + tokenizer = fit_tokenizer(args, splits["train"]) + cache_root = Path( + getattr(args, "paper_cache_root", None) + or (Path(args.data_root) / ".graph_tokenizer_cache") + ) + encoded_cache_key = _encoded_splits_cache_key( + tokenizer, splits, int(preset["max_position_embeddings"])) + encoded_cache_path = ( + cache_root / spec.canonical_name / "encoded" + / f"{encoded_cache_key}.pkl") + encoded_cache_hit = encoded_cache_path.is_file() + encoded, effective_max_length = _load_or_encode_splits( + tokenizer, + splits, + int(preset["max_position_embeddings"]), + encoded_cache_path, + encoded_cache_key, + ) + train_max_length = max( + (len(sequence) for sequence in encoded["train"]["input_ids"]), + default=1, + ) + normalizer, output_dim = _prepare_labels( + encoded, spec, getattr(args, "target_property", None)) + metric_semantics = _metric_semantics(spec, normalizer) + model_vocab_size = tokenizer.max_token_id + 1 + tokenizer.validate_model_vocab(model_vocab_size) + pooling = getattr(args, "pooling", "mean") + device = torch.device(getattr(args, "device", "cuda")) + precision = _resolve_precision_options(torch, device, preset) + if device.type == "cuda": + torch.backends.cuda.matmul.allow_tf32 = precision["allow_tf32"] + torch.backends.cudnn.allow_tf32 = precision["allow_tf32"] + if device.type == "cuda": + memory_device = _cuda_memory_query_device(torch, device) + try: + available_bytes, _ = torch.cuda.mem_get_info(memory_device) + except (AttributeError, TypeError): + available_bytes = torch.cuda.get_device_properties( + memory_device).total_memory + else: + available_bytes = 1 << 60 + memory_plan = _resolve_memory_plan( + encoder_type=preset["encoder"], + effective_batch_size=int(preset["batch_size"]), + sequence_length=train_max_length, + available_bytes=available_bytes, + activation_bytes_per_element=( + 3 if precision["torch_dtype"] is not None else 4), + ) + micro_batch_size = int(memory_plan["micro_batch_size"]) + accumulation_steps = int(memory_plan["gradient_accumulation_steps"]) + training_loss_name = str(preset.get("training_loss", "default")) + molhiv_pos_weight = None + if ( + spec.canonical_name == "molhiv" + and bool(preset.get("molhiv_pos_weight", False))): + molhiv_pos_weight = _molhiv_positive_weight( + torch, encoded["train"], device) + run_summaries = [] + experiment_fingerprints = [] + for run_index, seed in enumerate(seeds): + set_paper_seed(seed, torch) + experiment_fingerprint = _experiment_fingerprint( + preset, + spec, + encoded_cache_key, + model_vocab_size, + pooling, + seed=seed, + run_index=run_index, + ) + experiment_fingerprints.append(experiment_fingerprint) + grad_scaler = _new_grad_scaler(torch, precision["amp_dtype"]) + model = _create_paper_model( + encoder_type=preset["encoder"], + vocab_size=model_vocab_size, + pad_token_id=tokenizer.special_tokens.pad_token_id, + task_type=spec.task_type, + output_dim=output_dim, + pooling=pooling, + model_config={ + "max_position_embeddings": int(preset["max_position_embeddings"]), + }, + allow_random_gte_init=bool( + getattr(args, "allow_random_gte_init", False)), + pretrained_cache_dir=getattr(args, "gte_cache_dir", None), + ).to(device) + _require_fp32_parameters(torch, model) + loader_options = { + "pad_token_id": tokenizer.special_tokens.pad_token_id, + "num_workers": getattr(args, "num_workers", 0), + "pin_memory": device.type == "cuda", + } + train_loader = _make_loader( + torch, encoded["train"], micro_batch_size, shuffle=True, + bucket_by_length=True, **loader_options) + val_loader = _make_loader( + torch, encoded["val"], micro_batch_size, shuffle=False, + bucket_by_length=False, **loader_options) + + optimizer_steps = _optimizer_steps_per_epoch( + train_loader, accumulation_steps) + run_directory = ( + Path(args.output_dir) / "paper" / spec.canonical_name + / preset["encoder"] / f"run_{run_index}") + checkpoint_path = run_directory / "best.pt" + resume_path = run_directory / "last_state.pt" + resume_interval = max(1, int(preset.get("resume_interval", 5))) + resume_state = None + resume_phase = None + if bool(getattr(args, "resume", False)) and resume_path.is_file(): + resume_state = _torch_load(torch, resume_path, device) + if resume_state.get("checkpoint_kind") != "resume_state": + raise ValueError(f"Not a resume-state checkpoint: {resume_path}") + _validate_experiment_fingerprint( + resume_state, experiment_fingerprint) + resume_phase = resume_state.get("phase") + model.load_state_dict(resume_state["model"]) + + pretrain_epochs = int(preset["pretrain_epochs"]) + if resume_phase not in { + "pretrain_complete", "finetune", "finetune_complete"}: + pretrain_optimizer = torch.optim.AdamW( + model.parameters(), + lr=float(preset["pretrain_lr"]), + weight_decay=float(preset["weight_decay"]), + ) + pretrain_scheduler = _linear_warmup_scheduler( + torch, pretrain_optimizer, + optimizer_steps * pretrain_epochs, + preset["pretrain_warmup_ratio"]) + pretrain_start = 1 + if resume_phase == "pretrain": + resumed = restore_paper_resume_state( + resume_path, torch, model, pretrain_optimizer, + pretrain_scheduler, device, + grad_scaler=grad_scaler, + expected_fingerprint=experiment_fingerprint) + pretrain_start = int(resumed["epoch"]) + 1 + for epoch in range(pretrain_start, pretrain_epochs + 1): + if device.type == "cuda": + torch.cuda.reset_peak_memory_stats(device) + torch.cuda.synchronize(device) + started = time.monotonic() + mlm_loss = _train_mlm( + torch, model, train_loader, pretrain_optimizer, + pretrain_scheduler, tokenizer, device, + preset["pretrain_max_grad_norm"], preset["mask_prob"], + gradient_accumulation_steps=accumulation_steps, + torch_dtype=precision["torch_dtype"], + grad_scaler=grad_scaler) + if device.type == "cuda": + torch.cuda.synchronize(device) + runtime_metrics = _epoch_runtime_metrics( + torch, device, len(encoded["train"]), + time.monotonic() - started) + print(json.dumps({ + "event": "paper_epoch", + "phase": "pretrain", + "dataset": spec.canonical_name, + "model": preset["encoder"], + "run": run_index, + "seed": seed, + "epoch": epoch, + "epochs": pretrain_epochs, + "loss": mlm_loss, + **runtime_metrics, + }), flush=True) + if epoch % resume_interval == 0 or epoch == pretrain_epochs: + save_paper_resume_state( + resume_path, torch, model, pretrain_optimizer, + pretrain_scheduler, phase="pretrain", epoch=epoch, + grad_scaler=grad_scaler, + experiment_fingerprint=experiment_fingerprint) + save_paper_resume_state( + resume_path, torch, model, pretrain_optimizer, + pretrain_scheduler, phase="pretrain_complete", + epoch=pretrain_epochs, + grad_scaler=grad_scaler, + experiment_fingerprint=experiment_fingerprint) + elif resume_state is not None: + _restore_rng_state(torch, resume_state) + _restore_grad_scaler_state(resume_state, grad_scaler) + + optimizer = torch.optim.AdamW( + model.parameters(), + lr=float(preset["finetune_lr"]), + weight_decay=float(preset["weight_decay"]), + ) + finetune_epochs = int(preset["finetune_epochs"]) + scheduler = _linear_warmup_scheduler( + torch, optimizer, + optimizer_steps * finetune_epochs, + preset["finetune_warmup_ratio"]) + history = [] + best_metric = None + stale_epochs = 0 + finetune_start = 1 + finetune_checkpoint_complete = resume_phase == "finetune_complete" + skip_finetuning = finetune_checkpoint_complete + if resume_phase in {"finetune", "finetune_complete"}: + resumed = restore_paper_resume_state( + resume_path, torch, model, optimizer, scheduler, device, + grad_scaler=grad_scaler, + expected_fingerprint=experiment_fingerprint) + finetune_start = int(resumed["epoch"]) + 1 + extra = dict(resumed.get("extra", {})) + history = list(extra.get("history", [])) + best_metric = extra.get("best_metric") + stale_epochs = int(extra.get("stale_epochs", 0)) + skip_finetuning = _should_skip_finetuning( + resume_phase, stale_epochs, int(preset["patience"])) + if skip_finetuning: + finetune_start = finetune_epochs + 1 + higher_is_better = spec.canonical_name == "molhiv" + final_training_epoch = ( + int(resumed["epoch"]) + if skip_finetuning else finetune_start - 1) + for epoch in range(finetune_start, finetune_epochs + 1): + if device.type == "cuda": + torch.cuda.reset_peak_memory_stats(device) + torch.cuda.synchronize(device) + started = time.monotonic() + training_loss = _train_supervised( + torch, model, train_loader, optimizer, scheduler, spec, device, + preset["finetune_max_grad_norm"], + gradient_accumulation_steps=accumulation_steps, + loss_name=training_loss_name, + pos_weight=molhiv_pos_weight, + torch_dtype=precision["torch_dtype"], + grad_scaler=grad_scaler) + validation = _evaluate( + torch, model, val_loader, spec, device, normalizer, + torch_dtype=precision["torch_dtype"]) + if device.type == "cuda": + torch.cuda.synchronize(device) + runtime_metrics = _epoch_runtime_metrics( + torch, device, len(encoded["train"]), + time.monotonic() - started) + history.append({"epoch": epoch, "val_metric": validation["metric"], "val_loss": validation["loss"]}) + final_training_epoch = epoch + improved = best_metric is None or ( + validation["metric"] > best_metric + if higher_is_better else validation["metric"] < best_metric + ) + if improved: + best_metric = validation["metric"] + stale_epochs = 0 + save_paper_checkpoint( + checkpoint_path, torch, model, epoch, best_metric, + normalizer, + experiment_fingerprint=experiment_fingerprint) + else: + stale_epochs += 1 + print(json.dumps({ + "event": "paper_epoch", + "phase": "finetune", + "dataset": spec.canonical_name, + "model": preset["encoder"], + "run": run_index, + "seed": seed, + "epoch": epoch, + "epochs": finetune_epochs, + "training_loss": training_loss, + "val_loss": validation["loss"], + "val_metric": validation["metric"], + "best_metric": best_metric, + "stale_epochs": stale_epochs, + **runtime_metrics, + }), flush=True) + should_stop = stale_epochs >= int(preset["patience"]) + if ( + epoch % resume_interval == 0 + or epoch == finetune_epochs + or should_stop): + save_paper_resume_state( + resume_path, torch, model, optimizer, scheduler, + phase="finetune", epoch=epoch, extra={ + "history": history, + "best_metric": best_metric, + "stale_epochs": stale_epochs, + "normalizer": normalizer, + }, grad_scaler=grad_scaler, + experiment_fingerprint=experiment_fingerprint) + if should_stop: + break + if not finetune_checkpoint_complete: + save_paper_resume_state( + resume_path, torch, model, optimizer, scheduler, + phase="finetune_complete", epoch=final_training_epoch, extra={ + "history": history, + "best_metric": best_metric, + "stale_epochs": stale_epochs, + "normalizer": normalizer, + }, grad_scaler=grad_scaler, + experiment_fingerprint=experiment_fingerprint) + state = restore_paper_checkpoint( + checkpoint_path, torch, model, device, + expected_fingerprint=experiment_fingerprint) + final_validation = _evaluate( + torch, model, val_loader, spec, device, state["normalizer"], + torch_dtype=precision["torch_dtype"]) + test_loader = _make_loader( + torch, encoded["test"], micro_batch_size, shuffle=False, + bucket_by_length=False, **loader_options) + final_test = _evaluate( + torch, model, test_loader, spec, device, state["normalizer"], + torch_dtype=precision["torch_dtype"]) + run_summaries.append({ + "run": run_index, + "seed": seed, + "best_epoch": state["epoch"], + "best_val": final_validation, + "best_test": final_test, + "checkpoint_path": str(checkpoint_path), + "resume_checkpoint_path": str(resume_path), + "resumed": resume_state is not None, + "history": history, + "test_evaluations": 1, + "experiment_fingerprint": experiment_fingerprint, + "model": model.manifest(), + **metric_semantics, + }) + return { + "dataset": spec.canonical_name, + "model": preset["encoder"], + "protocol": "paper", + "runs": run_summaries, + "preset": preset, + "effective_max_length": effective_max_length, + "tokenizer_vocab_size": tokenizer.max_token_id + 1, + "experiment_fingerprints": experiment_fingerprints, + "memory_plan": memory_plan, + "precision": { + "amp_dtype": precision["amp_dtype"], + "allow_tf32": precision["allow_tf32"], + }, + "training_loss": { + "name": training_loss_name, + "molhiv_pos_weight": ( + float(molhiv_pos_weight.item()) + if molhiv_pos_weight is not None else None), + }, + **metric_semantics, + "cache": { + "tokenizer_status": getattr(tokenizer, "_cache_status", "unknown"), + "tokenizer_path": getattr(tokenizer, "_cache_path", None), + "encoded_status": "hit" if encoded_cache_hit else "miss", + "encoded_path": str(encoded_cache_path), + }, + } diff --git a/examples/graph_tokenizer/requirements.txt b/examples/graph_tokenizer/requirements.txt new file mode 100644 index 000000000..84e11195e --- /dev/null +++ b/examples/graph_tokenizer/requirements.txt @@ -0,0 +1,8 @@ +# GraphTokenizer paper-mode dependencies; install GammaGL itself separately. +# The paper protocol verifies the exact runtime versions before a formal run. +torch==2.1.2 +dgl==2.4.0 +torch-geometric==2.4.0 +ogb>=1.3.6 +huggingface-hub>=0.20 +safetensors>=0.4 diff --git a/gammagl/datasets/__init__.py b/gammagl/datasets/__init__.py index 1c3351876..cd3391603 100644 --- a/gammagl/datasets/__init__.py +++ b/gammagl/datasets/__init__.py @@ -20,6 +20,9 @@ from .wikics import WikiCS from .blogcatalog import BlogCatalog from .molecule_net import MoleculeNet +from .qm9 import QM9 +from .ogbg_molhiv import OGBGMolHIV +from .peptides_struct import PeptidesStruct from .facebook import FacebookPagePage from .acm4heco import ACM4HeCo from .yelp import Yelp @@ -53,6 +56,9 @@ 'PolBlogs', 'WikiCS', 'MoleculeNet', + 'QM9', + 'OGBGMolHIV', + 'PeptidesStruct', 'FacebookPagePage', 'NGSIM_US_101', 'Yelp', diff --git a/gammagl/datasets/_graph_tokenizer_download.py b/gammagl/datasets/_graph_tokenizer_download.py new file mode 100644 index 000000000..98fa5ccf8 --- /dev/null +++ b/gammagl/datasets/_graph_tokenizer_download.py @@ -0,0 +1,228 @@ +import hashlib +import os +import shutil +import tarfile +import tempfile +import zipfile +from contextlib import contextmanager +from pathlib import Path +from typing import Iterable, Tuple + + +DATA_BUNDLE_ENV = 'GAMMAGL_GRAPH_TOKENIZER_DATA_BUNDLE' +PAPER_DATA_BUNDLE_ID = '10etZF9OnV569_Fp7tpdMUVEH9eZECKdW' +PAPER_DATA_BUNDLE_FILENAME = 'GraphToenizerDataset.tar.gz' +# SHA-256 of the bytes published by the fixed Google Drive release above. +PAPER_DATA_BUNDLE_SHA256 = ( + '5c437c3c0d4b7278379c0e70d57f98148e5c815d753d8cf68e2a45952bcce459') +_REQUIRED_SPLITS = ('train_index.json', 'val_index.json', 'test_index.json') +_OPTIONAL_FILES = ('token_mappings.json',) + + +def materialize_paper_dataset( + dataset_name: str, + aliases: Iterable[str], + raw_dir, + cache_root, + allow_download: bool = False, +) -> Path: + """Copy one dataset from a verified, already-prepared paper bundle.""" + raw_dir = Path(raw_dir) + cache_root = Path(cache_root) + source = _resolve_bundle_source(cache_root, allow_download=allow_download) + bundle_root = source if source.is_dir() else _extract_archive(source, cache_root) + dataset_dir = _find_dataset_dir(bundle_root, dataset_name, tuple(aliases)) + + data_file = _find_data_file(dataset_dir) + missing_splits = [name for name in _REQUIRED_SPLITS if not (dataset_dir / name).is_file()] + if data_file is None or missing_splits: + details = [] + if data_file is None: + details.append('data.pkl or data.pkl.gz') + if missing_splits: + details.append(f"official split files: {', '.join(missing_splits)}") + raise FileNotFoundError( + f"Released {dataset_name} directory is incomplete; missing {'; '.join(details)}.") + + raw_dir.mkdir(parents=True, exist_ok=True) + for source_file in (data_file, *(dataset_dir / name for name in _REQUIRED_SPLITS)): + shutil.copy2(source_file, raw_dir / source_file.name) + for name in _OPTIONAL_FILES: + source_file = dataset_dir / name + if source_file.is_file(): + shutil.copy2(source_file, raw_dir / name) + return raw_dir + + +@contextmanager +def _bundle_lock(cache_dir: Path): + """Serialize installation of the shared release archive.""" + try: + import fcntl + except ImportError as error: # pragma: no cover - paper runs are Linux-only. + raise RuntimeError("GraphTokenizer paper preparation requires fcntl locking.") from error + lock_path = cache_dir / ".bundle.lock" + with lock_path.open("a+") as handle: + fcntl.flock(handle.fileno(), fcntl.LOCK_EX) + try: + yield + finally: + fcntl.flock(handle.fileno(), fcntl.LOCK_UN) + + +def _resolve_bundle_source(cache_root: Path, allow_download: bool = False) -> Path: + configured = os.environ.get(DATA_BUNDLE_ENV) + if configured: + if not allow_download: + raise FileNotFoundError( + "GraphTokenizer paper data must be prepared by the single " + "data-preparation process before starting training workers.") + source = Path(configured).expanduser().resolve() + if not source.exists(): + raise FileNotFoundError( + f"{DATA_BUNDLE_ENV} points to a missing path: {source}") + if source.is_file(): + _verify_official_bundle(source, remove_on_failure=False) + return source + + cache_dir = cache_root / '.graph_tokenizer_release' + cache_dir.mkdir(parents=True, exist_ok=True) + archive = cache_dir / PAPER_DATA_BUNDLE_FILENAME + with _bundle_lock(cache_dir): + if archive.exists(): + _verify_official_bundle(archive, remove_on_failure=False) + elif not allow_download: + raise FileNotFoundError( + "GraphTokenizer paper data must be prepared by the single " + "data-preparation process before starting training workers.") + else: + from gammagl.data.download import download_google_url + + descriptor, temporary_name = tempfile.mkstemp( + prefix=f".{archive.name}.{os.getpid()}.", suffix=".tmp", dir=cache_dir) + os.close(descriptor) + temporary = Path(temporary_name) + try: + download_google_url( + PAPER_DATA_BUNDLE_ID, + str(cache_dir), + temporary.name, + ) + _verify_official_bundle(temporary, remove_on_failure=False) + os.replace(temporary, archive) + except Exception: + if temporary.exists(): + temporary.unlink() + raise + return archive + + +def sha256_file(path: Path, chunk_size: int = 1024 * 1024) -> str: + """Return the streaming SHA-256 digest of a file's actual bytes.""" + digest = hashlib.sha256() + with Path(path).open('rb') as handle: + for chunk in iter(lambda: handle.read(chunk_size), b''): + digest.update(chunk) + return digest.hexdigest() + + +def _verify_official_bundle(archive: Path, remove_on_failure: bool) -> str: + actual = sha256_file(archive) + if actual == PAPER_DATA_BUNDLE_SHA256: + return actual + if remove_on_failure and archive.exists(): + archive.unlink() + raise RuntimeError( + 'GraphTokenizer paper bundle SHA-256 mismatch: expected ' + f'{PAPER_DATA_BUNDLE_SHA256}, got {actual}.') + + +def _extract_archive(archive: Path, cache_root: Path) -> Path: + digest = _verify_official_bundle(archive, remove_on_failure=False)[:16] + destination = cache_root / '.graph_tokenizer_release' / f'extracted-{digest}' + marker = destination / '.complete' + if marker.is_file(): + return destination + + destination.mkdir(parents=True, exist_ok=True) + if zipfile.is_zipfile(archive): + _extract_zip(archive, destination) + elif tarfile.is_tarfile(archive): + _extract_tar(archive, destination) + else: + raise ValueError( + f'Unsupported paper data bundle format: {archive}. Expected ZIP or TAR.') + marker.touch() + return destination + + +def _safe_member_path(destination: Path, member_name: str) -> Path: + member_path = (destination / member_name).resolve() + try: + member_path.relative_to(destination.resolve()) + except ValueError as error: + raise ValueError(f'Unsafe path in paper data bundle: {member_name}') from error + return member_path + + +def _extract_zip(archive: Path, destination: Path) -> None: + with zipfile.ZipFile(archive) as handle: + for member in handle.infolist(): + target = _safe_member_path(destination, member.filename) + if member.is_dir(): + target.mkdir(parents=True, exist_ok=True) + continue + target.parent.mkdir(parents=True, exist_ok=True) + with handle.open(member) as source, target.open('wb') as output: + shutil.copyfileobj(source, output) + + +def _extract_tar(archive: Path, destination: Path) -> None: + with tarfile.open(archive) as handle: + for member in handle.getmembers(): + if member.issym() or member.islnk(): + raise ValueError(f'Links are not allowed in paper data bundle: {member.name}') + target = _safe_member_path(destination, member.name) + if member.isdir(): + target.mkdir(parents=True, exist_ok=True) + continue + if not member.isfile(): + continue + target.parent.mkdir(parents=True, exist_ok=True) + source = handle.extractfile(member) + if source is None: + raise ValueError(f'Cannot extract paper data bundle member: {member.name}') + with source, target.open('wb') as output: + shutil.copyfileobj(source, output) + + +def _find_dataset_dir( + bundle_root: Path, + dataset_name: str, + aliases: Tuple[str, ...], +) -> Path: + accepted_names = {_normalize_name(dataset_name)} + accepted_names.update(_normalize_name(alias) for alias in aliases) + candidates = [] + for data_file in (*bundle_root.rglob('data.pkl'), *bundle_root.rglob('data.pkl.gz')): + parent = data_file.parent + if _normalize_name(parent.name) in accepted_names: + candidates.append(parent) + unique_candidates = sorted(set(candidates), key=lambda path: (len(path.parts), str(path))) + if not unique_candidates: + names = ', '.join(sorted(accepted_names)) + raise FileNotFoundError( + f"Cannot find {dataset_name} ({names}) in released paper data bundle {bundle_root}.") + return unique_candidates[0] + + +def _find_data_file(dataset_dir: Path): + for name in ('data.pkl', 'data.pkl.gz'): + path = dataset_dir / name + if path.is_file(): + return path + return None + + +def _normalize_name(value: str) -> str: + return str(value).strip().lower().replace('_', '-') diff --git a/gammagl/datasets/_molecular_benchmark.py b/gammagl/datasets/_molecular_benchmark.py new file mode 100644 index 000000000..573e5e89f --- /dev/null +++ b/gammagl/datasets/_molecular_benchmark.py @@ -0,0 +1,369 @@ +import gzip +import json +import math +import os +import os.path as osp +import pickle +from collections.abc import Iterable +from typing import Callable, Dict, List, Optional, Sequence, Tuple + +import tensorlayerx as tlx + +from gammagl.data import Graph, InMemoryDataset, download_url +from gammagl.data.extract import extract_tar, extract_zip +from gammagl.data.makedirs import makedirs + +from ._graph_tokenizer_download import materialize_paper_dataset + + +def validate_molecular_splits(num_samples: int, split_indices: Dict[str, List[int]]) -> None: + required = ('train', 'val', 'test') + if set(split_indices) != set(required): + raise ValueError(f"Split files must define exactly {required}.") + + seen = {} + for split in required: + for index in split_indices[split]: + if index < 0 or index >= num_samples: + raise ValueError( + f"{split} split index {index} is outside [0, {num_samples}).") + if index in seen: + raise ValueError( + f"{split} split index {index} overlaps {seen[index]} split.") + seen[index] = split + + +class PreprocessedMolecularBenchmark(InMemoryDataset): + r"""Base class for molecular benchmark datasets stored as preprocessed + graph pickle files plus train/validation/test split indices. + + The expected raw layout is: + + .. code-block:: text + + root/name/raw/data.pkl or data.pkl.gz + root/name/raw/train_index.json + root/name/raw/val_index.json + root/name/raw/test_index.json + """ + + name: str = "" + display_name: str = "" + url: Optional[str] = None + aliases: Tuple[str, ...] = () + metric: str = "" + num_tasks: int = 1 + task_type: str = "" + label_keys: Tuple[str, ...] = () + allow_nan_labels: bool = False + node_feature_columns: Dict[str, int] = {} + edge_feature_columns: Dict[str, int] = {} + + def __init__( + self, + root: str, + transform: Optional[Callable] = None, + pre_transform: Optional[Callable] = None, + pre_filter: Optional[Callable] = None, + force_reload: bool = False, + ) -> None: + super().__init__(root, transform, pre_transform, pre_filter, + force_reload=force_reload) + self._download() + self._process() + self.data, self.slices = self.load_data(self.processed_paths[0]) + self.split_indices = self._load_processed_splits() + + @property + def raw_file_names(self) -> List[str]: + return ['train_index.json', 'val_index.json', 'test_index.json'] + + @property + def processed_file_names(self) -> List[str]: + return [tlx.BACKEND + '_data.pt', 'split_indices.json'] + + def _download(self): + if self._has_raw_files() or self._has_processed_files(): + return + makedirs(self.raw_dir) + self.download() + + def _has_raw_files(self) -> bool: + return self._raw_data_file() is not None and all( + osp.exists(osp.join(self.raw_dir, name)) for name in self.raw_file_names) + + def _has_processed_files(self) -> bool: + return all(osp.exists(path) for path in self.processed_paths) + + def _raw_data_file(self) -> Optional[str]: + for filename in ('data.pkl', 'data.pkl.gz'): + path = osp.join(self.raw_dir, filename) + if osp.exists(path): + return path + return None + + def download(self) -> None: + if self.url is None: + materialize_paper_dataset( + dataset_name=self.name, + aliases=self.aliases, + raw_dir=self.raw_dir, + cache_root=self.root, + ) + return + + path = download_url(self.url, self.raw_dir) + if path.endswith('.zip'): + extract_zip(path, self.raw_dir) + os.unlink(path) + elif path.endswith(('.tar.gz', '.tgz')): + extract_tar(path, self.raw_dir) + os.unlink(path) + + def process(self) -> None: + data_file = self._raw_data_file() + if data_file is None: + raise FileNotFoundError( + f"Expected data.pkl or data.pkl.gz under {self.raw_dir}") + + raw_samples = self._read_pickle(data_file) + split_indices = self._read_split_indices() + validate_molecular_splits(len(raw_samples), split_indices) + data_list = [ + self._sample_to_graph(sample, index) + for index, sample in enumerate(raw_samples) + ] + + if self.pre_filter is not None: + raise ValueError( + "pre_filter is not supported for benchmark datasets because it " + "would invalidate the fixed official split indices.") + + if self.pre_transform is not None: + data_list = [self.pre_transform(data) for data in data_list] + + self.save_data(self.collate(data_list), self.processed_paths[0]) + with open(self.processed_paths[1], 'w', encoding='utf-8') as f: + json.dump(split_indices, f) + + def get_idx_split(self) -> Dict[str, List[int]]: + return {key: list(value) for key, value in self.split_indices.items()} + + def get_split(self, split: str): + if split not in self.split_indices: + expected = ', '.join(sorted(self.split_indices)) + raise ValueError(f"Unknown split '{split}'. Expected one of: {expected}.") + return self.index_select(self.split_indices[split]) + + def _load_processed_splits(self) -> Dict[str, List[int]]: + with open(self.processed_paths[1], 'r', encoding='utf-8') as f: + split_indices = {key: [int(index) for index in value] + for key, value in json.load(f).items()} + validate_molecular_splits(len(self), split_indices) + return split_indices + + def _read_pickle(self, path: str): + opener = gzip.open if path.endswith('.gz') else open + with opener(path, 'rb') as f: + return pickle.load(f) + + def _read_split_indices(self) -> Dict[str, List[int]]: + split_indices = {} + for split in ('train', 'val', 'test'): + path = osp.join(self.raw_dir, f'{split}_index.json') + if not osp.exists(path): + raise FileNotFoundError(f"Split file not found: {path}") + with open(path, 'r', encoding='utf-8') as f: + split_indices[split] = [int(index) for index in json.load(f)] + return split_indices + + def _sample_to_graph(self, sample, index: int) -> Graph: + if isinstance(sample, Graph): + graph = self._graph_from_mapping({ + 'edge_index': getattr(sample, 'edge_index', None), + 'x': getattr(sample, 'x', None), + 'edge_attr': getattr(sample, 'edge_attr', None), + }) + label_data = getattr(sample, 'y', None) + elif isinstance(sample, dict): + graph = self._graph_from_mapping(sample) + label_data = sample.get('properties', sample.get('y', sample.get('label', sample.get('labels')))) + elif isinstance(sample, tuple) and len(sample) >= 2: + graph = self._graph_from_mapping(sample[0]) + label_data = sample[1] + else: + raise ValueError(f"Unsupported sample format at index {index}: {type(sample)!r}") + + graph.y = tlx.convert_to_tensor([self._extract_label(label_data)], dtype=tlx.float32) + return graph + + def _graph_from_mapping(self, graph_data) -> Graph: + if isinstance(graph_data, Graph): + return graph_data + if not isinstance(graph_data, dict) and hasattr(graph_data, 'edges'): + src, dst = graph_data.edges() + node_data = getattr(graph_data, 'ndata', {}) + edge_data = getattr(graph_data, 'edata', {}) + node_name, node_values = self._first_present( + node_data, ('node_token_ids', 'x', 'attr')) + edge_name, edge_values = self._first_present( + edge_data, ('edge_token_ids', 'edge_attr')) + graph_data = { + 'edge_index': [self._to_list(src), self._to_list(dst)], + node_name: node_values, + edge_name: edge_values, + } + graph_data = {key: value for key, value in graph_data.items() + if key is not None and value is not None} + if not isinstance(graph_data, dict): + raise ValueError(f"Unsupported graph object: {type(graph_data)!r}") + + edge_index = graph_data.get('edge_index') + if edge_index is None and 'edges' in graph_data: + edge_index = graph_data['edges'] + + x_name, x = self._first_present( + graph_data, + ('node_token_ids', 'node_type_ids', 'x', 'node_features', 'node_feat', + 'attr', 'node_labels')) + edge_name, edge_attr = self._first_present( + graph_data, + ('edge_token_ids', 'edge_type_ids', 'edge_attr', 'edge_features', + 'edge_feat', 'edge_labels')) + if x is None: + raise ValueError("Graph is missing paper-required node features.") + if edge_attr is None: + raise ValueError("Graph is missing paper-required edge features.") + + return Graph( + x=tlx.convert_to_tensor( + self._encode_feature_ids(x_name, x, is_edge=False), dtype=tlx.int64), + edge_index=tlx.convert_to_tensor(self._normalize_edge_index(edge_index), dtype=tlx.int64), + edge_attr=tlx.convert_to_tensor( + self._encode_feature_ids(edge_name, edge_attr, is_edge=True), dtype=tlx.int64), + y=None, + ) + + def _extract_label(self, label_data) -> List[float]: + if isinstance(label_data, dict): + if self.label_keys == ('labels',): + if 'labels' not in label_data: + raise ValueError("Label mapping is missing required 'labels' field.") + return self._as_float_list(label_data['labels']) + missing = [key for key in self.label_keys if key not in label_data] + if missing: + if self.name == 'qm9': + raise ValueError( + "QM9 sample must provide exactly 16 QM9 properties; " + f"missing: {', '.join(missing)}.") + raise ValueError(f"Label mapping is missing required fields: {', '.join(missing)}.") + return self._as_float_list([label_data[key] for key in self.label_keys]) + return self._as_float_list(label_data) + + def _as_float_list(self, value) -> List[float]: + if value is None: + raise ValueError(f"{self.display_name} sample is missing its label.") + values = self._to_list(value) + while self._is_sequence(values) and len(values) == 1 and self._is_sequence(values[0]): + values = values[0] + if not self._is_sequence(values): + values = [values] + if len(values) != self.num_tasks: + raise ValueError( + f"{self.display_name} labels must have exactly {self.num_tasks} values; " + f"received {len(values)}.") + result = [float(item) for item in values] + if not self.allow_nan_labels and any(math.isnan(item) for item in result): + raise ValueError(f"{self.display_name} labels cannot contain NaN values.") + return result + + @staticmethod + def _first_present(mapping, names: Sequence[str]): + for name in names: + if name in mapping: + return name, mapping[name] + return None, None + + @classmethod + def _normalize_edge_index(cls, edge_index) -> List[List[int]]: + if edge_index is None: + return [[], []] + values = cls._to_list(edge_index) + if len(values) == 2 and cls._is_sequence(values[0]) and cls._is_sequence(values[1]): + return [[int(value) for value in values[0]], [int(value) for value in values[1]]] + src = [int(pair[0]) for pair in values] + dst = [int(pair[1]) for pair in values] + return [src, dst] + + def _encode_feature_ids(self, field_name: str, values, is_edge: bool) -> List[int]: + values = self._to_list(values) + if not self._is_sequence(values): + raise ValueError(f"{field_name} must be a sequence of feature values.") + + if field_name.endswith('_token_ids'): + return [self._scalar_token_id(field_name, item) for item in values] + + encoded = [] + for position, item in enumerate(values): + raw_type = self._extract_feature_type(field_name, item, is_edge) + if raw_type < 0: + raise ValueError( + f"{field_name}[{position}] must be non-negative; received {raw_type}.") + encoded.append(2 * raw_type if is_edge else 2 * raw_type + 1) + return encoded + + def _scalar_token_id(self, field_name: str, item) -> int: + item = self._to_list(item) + if self._is_sequence(item): + if len(item) != 1 or self._is_sequence(item[0]): + raise ValueError( + f"{field_name} must contain one scalar token per node or edge.") + item = item[0] + token_id = int(item) + if token_id < 0: + raise ValueError(f"{field_name} token IDs must be non-negative.") + return token_id + + def _extract_feature_type(self, field_name: str, item, is_edge: bool) -> int: + item = self._to_list(item) + if not self._is_sequence(item): + return int(item) + if not item: + raise ValueError(f"{field_name} contains an empty feature row.") + + if is_edge and self._is_one_hot(item): + return int(max(range(len(item)), key=lambda index: item[index])) + 1 + + columns = self.edge_feature_columns if is_edge else self.node_feature_columns + column = columns.get(field_name, 0) + if column >= len(item): + raise ValueError( + f"{field_name} requires feature column {column}, but row has {len(item)} values.") + return int(item[column]) + + @staticmethod + def _is_one_hot(values) -> bool: + try: + numeric = [float(value) for value in values] + except (TypeError, ValueError): + return False + return all(value in (0.0, 1.0) for value in numeric) and sum(numeric) == 1.0 + + @staticmethod + def _to_list(value): + if hasattr(value, 'detach'): + value = value.detach() + if hasattr(value, 'cpu'): + value = value.cpu() + if hasattr(value, 'numpy'): + value = value.numpy() + if hasattr(value, 'tolist'): + return value.tolist() + return value + + @staticmethod + def _is_sequence(value) -> bool: + return isinstance(value, Iterable) and not isinstance(value, (str, bytes, dict)) + + def __repr__(self) -> str: + return f'{self.display_name}({len(self)})' diff --git a/gammagl/datasets/ogbg_molhiv.py b/gammagl/datasets/ogbg_molhiv.py new file mode 100644 index 000000000..c46666d69 --- /dev/null +++ b/gammagl/datasets/ogbg_molhiv.py @@ -0,0 +1,17 @@ +from ._molecular_benchmark import PreprocessedMolecularBenchmark + + +class OGBGMolHIV(PreprocessedMolecularBenchmark): + r"""The OGBG-molhiv molecular property prediction benchmark. + + Data must be prepared by the GraphTokenizer single-process preparation + command before construction; training never downloads the shared bundle. + """ + + name = 'ogbg-molhiv' + display_name = 'OGBG-molhiv' + aliases = ('molhiv', 'ogbg-molhiv', 'ogbg_molhiv') + task_type = 'binary_classification' + num_tasks = 1 + metric = 'rocauc' + label_keys = ('label',) diff --git a/gammagl/datasets/peptides_struct.py b/gammagl/datasets/peptides_struct.py new file mode 100644 index 000000000..222c235f6 --- /dev/null +++ b/gammagl/datasets/peptides_struct.py @@ -0,0 +1,18 @@ +from ._molecular_benchmark import PreprocessedMolecularBenchmark + + +class PeptidesStruct(PreprocessedMolecularBenchmark): + r"""The Peptides-struct molecular graph regression benchmark. + + Data must be prepared by the GraphTokenizer single-process preparation + command before construction; training never downloads the shared bundle. + """ + + name = 'peptides-struct' + display_name = 'Peptides-struct' + aliases = ('peptides-struct', 'peptides_struct', 'p-struct', 'p_struct') + task_type = 'multi_target_regression' + num_tasks = 11 + metric = 'average_mae' + label_keys = ('labels',) + allow_nan_labels = True diff --git a/gammagl/datasets/qm9.py b/gammagl/datasets/qm9.py new file mode 100644 index 000000000..943dca682 --- /dev/null +++ b/gammagl/datasets/qm9.py @@ -0,0 +1,23 @@ +from ._molecular_benchmark import PreprocessedMolecularBenchmark + + +class QM9(PreprocessedMolecularBenchmark): + r"""The QM9 molecular property prediction benchmark. + + This loader follows the GammaGL :class:`InMemoryDataset` workflow. On the + data is prepared by the GraphTokenizer single-process preparation command, + which caches the graph pickle and official split files under ``raw``. + """ + + name = 'qm9' + display_name = 'QM9' + aliases = ('qm9',) + task_type = 'regression' + num_tasks = 16 + metric = 'mae' + label_keys = ( + 'mu', 'alpha', 'homo', 'lumo', 'gap', 'r2', 'zpve', 'u0', + 'u298', 'h298', 'g298', 'cv', 'u0_atom', 'u298_atom', + 'h298_atom', 'g298_atom', + ) + node_feature_columns = {'attr': 5, 'x': 0} diff --git a/gammagl/models/__init__.py b/gammagl/models/__init__.py index 729eaabef..c67c925ac 100644 --- a/gammagl/models/__init__.py +++ b/gammagl/models/__init__.py @@ -72,6 +72,8 @@ from .amp import AMPModel, amp_elbo_regression_loss from .gnrf import GNRF, GNN from .defog import DeFoGModel +from .graph_bert import GraphBERT, GraphTokenTransformer +from .graph_gte import GraphGTE __all__ = [ 'HeCo', 'GCNModel', @@ -155,6 +157,9 @@ 'GNRF', 'GNN', 'DeFoGModel', + 'GraphTokenTransformer', + 'GraphBERT', + 'GraphGTE', ] classes = __all__ diff --git a/gammagl/models/graph_bert.py b/gammagl/models/graph_bert.py new file mode 100644 index 000000000..fc7e8574a --- /dev/null +++ b/gammagl/models/graph_bert.py @@ -0,0 +1,309 @@ +"""Native TensorLayerX BERT encoder for serialized graph tokens.""" + +from __future__ import annotations + +import math + +import tensorlayerx as tlx + + +def _truncated_normal(): + return tlx.initializers.TruncatedNormal(stddev=0.02) + + +def _linear(in_features: int, out_features: int, bias: bool = True): + return tlx.nn.Linear( + in_features=in_features, + out_features=out_features, + W_init=_truncated_normal(), + b_init=tlx.initializers.Constant(0.0) if bias else None, + ) + + +def _parameter_count(weights) -> int: + total = 0 + for weight in weights: + size = 1 + for dimension in tlx.get_tensor_shape(weight): + size *= int(dimension) + total += size + return total + + +def masked_mean(hidden_states, attention_mask=None): + if attention_mask is None: + return tlx.reduce_mean(hidden_states, axis=1) + mask = tlx.expand_dims(tlx.cast(attention_mask, tlx.float32), axis=-1) + numerator = tlx.reduce_sum(hidden_states * mask, axis=1) + denominator = tlx.reduce_sum(mask, axis=1) + denominator = tlx.maximum(denominator, tlx.ones_like(denominator)) + return numerator / denominator + + +def additive_attention_bias(attention_mask): + if attention_mask is None: + return None + mask = tlx.cast(attention_mask, tlx.float32) + mask = tlx.expand_dims(tlx.expand_dims(mask, axis=1), axis=1) + return (1.0 - mask) * -10000.0 + + +class GraphTokenSelfAttention(tlx.nn.Module): + """Backend-neutral scaled dot-product multi-head self-attention.""" + + def __init__( + self, + hidden_size: int, + num_attention_heads: int, + attention_dropout_rate: float = 0.1, + name: str | None = None, + ): + super().__init__(name=name) + if hidden_size % num_attention_heads != 0: + raise ValueError("hidden_size must be divisible by num_attention_heads.") + self.hidden_size = int(hidden_size) + self.num_attention_heads = int(num_attention_heads) + self.head_size = self.hidden_size // self.num_attention_heads + self.q_proj = _linear(self.hidden_size, self.hidden_size) + self.k_proj = _linear(self.hidden_size, self.hidden_size) + self.v_proj = _linear(self.hidden_size, self.hidden_size) + self.out_proj = _linear(self.hidden_size, self.hidden_size) + self.attention_dropout = tlx.nn.Dropout(p=float(attention_dropout_rate)) + + def _split_heads(self, value): + batch_size, sequence_length = tlx.get_tensor_shape(value)[:2] + value = tlx.reshape( + value, + [batch_size, sequence_length, self.num_attention_heads, self.head_size], + ) + return tlx.transpose(value, [0, 2, 1, 3]) + + def _merge_heads(self, value): + value = tlx.transpose(value, [0, 2, 1, 3]) + batch_size, sequence_length = tlx.get_tensor_shape(value)[:2] + return tlx.reshape(value, [batch_size, sequence_length, self.hidden_size]) + + def forward(self, hidden_states, attention_mask=None): + query = self._split_heads(self.q_proj(hidden_states)) + key = self._split_heads(self.k_proj(hidden_states)) + value = self._split_heads(self.v_proj(hidden_states)) + key_transposed = tlx.transpose(key, [0, 1, 3, 2]) + scores = tlx.matmul(query, key_transposed) / math.sqrt(self.head_size) + bias = additive_attention_bias(attention_mask) + if bias is not None: + scores = scores + bias + probabilities = tlx.softmax(scores, axis=-1) + probabilities = self.attention_dropout(probabilities) + context = tlx.matmul(probabilities, value) + return self.out_proj(self._merge_heads(context)) + + +class GraphTokenTransformerLayer(tlx.nn.Module): + """Post-normalization BERT encoder block.""" + + def __init__( + self, + hidden_size: int, + num_attention_heads: int, + intermediate_size: int, + dropout_rate: float = 0.1, + attention_dropout_rate: float | None = None, + layer_norm_eps: float = 1e-12, + name: str | None = None, + ): + super().__init__(name=name) + attention_dropout_rate = ( + dropout_rate if attention_dropout_rate is None + else attention_dropout_rate + ) + self.attention = GraphTokenSelfAttention( + hidden_size=hidden_size, + num_attention_heads=num_attention_heads, + attention_dropout_rate=attention_dropout_rate, + ) + self.attention_norm = tlx.nn.LayerNorm( + normalized_shape=hidden_size, + epsilon=layer_norm_eps, + gamma_init="ones", + beta_init="zeros", + ) + self.ffn_in = _linear(hidden_size, intermediate_size) + self.ffn_out = _linear(intermediate_size, hidden_size) + self.ffn_norm = tlx.nn.LayerNorm( + normalized_shape=hidden_size, + epsilon=layer_norm_eps, + gamma_init="ones", + beta_init="zeros", + ) + self.hidden_dropout = tlx.nn.Dropout(p=float(dropout_rate)) + + def forward(self, hidden_states, attention_mask=None): + attention_output = self.attention( + hidden_states, attention_mask=attention_mask) + hidden_states = self.attention_norm( + hidden_states + self.hidden_dropout(attention_output)) + ffn_output = self.ffn_out(tlx.gelu(self.ffn_in(hidden_states))) + return self.ffn_norm(hidden_states + self.hidden_dropout(ffn_output)) + + +class GraphTokenTransformer(tlx.nn.Module): + """BERT-Small encoder with MLM and graph-level task heads.""" + + encoder_type = "bert" + model_name = "bert-small" + position_embedding_type = "absolute" + hidden_act = "gelu" + + def __init__( + self, + vocab_size: int, + output_dim: int, + hidden_size: int = 512, + num_hidden_layers: int = 4, + num_attention_heads: int = 4, + intermediate_size: int = 2048, + max_position_embeddings: int = 768, + dropout_rate: float = 0.1, + attention_dropout_rate: float | None = None, + layer_norm_eps: float = 1e-12, + pad_token_id: int = 0, + type_vocab_size: int = 2, + task_type: str = "regression", + pooling: str = "mean", + task_dropout: float = 0.1, + name: str | None = None, + ): + super().__init__(name=name) + if pooling not in {"mean", "cls"}: + raise ValueError("pooling must be 'mean' or 'cls'.") + if max_position_embeddings <= 0: + raise ValueError("max_position_embeddings must be positive.") + attention_dropout_rate = ( + dropout_rate if attention_dropout_rate is None + else attention_dropout_rate + ) + self.vocab_size = int(vocab_size) + self.output_dim = int(output_dim) + self.hidden_size = int(hidden_size) + self.num_hidden_layers = int(num_hidden_layers) + self.num_attention_heads = int(num_attention_heads) + self.intermediate_size = int(intermediate_size) + self.max_position_embeddings = int(max_position_embeddings) + self.dropout_rate = float(dropout_rate) + self.attention_dropout_rate = float(attention_dropout_rate) + self.layer_norm_eps = float(layer_norm_eps) + self.pad_token_id = int(pad_token_id) + self.type_vocab_size = int(type_vocab_size) + self.task_type = str(task_type) + self.pooling = pooling + self.task_dropout_rate = float(task_dropout) + + self.token_embeddings = tlx.nn.Embedding( + num_embeddings=self.vocab_size, + embedding_dim=self.hidden_size, + E_init=_truncated_normal(), + ) + self.position_embeddings = tlx.nn.Embedding( + num_embeddings=self.max_position_embeddings, + embedding_dim=self.hidden_size, + E_init=_truncated_normal(), + ) + self.token_type_embeddings = tlx.nn.Embedding( + num_embeddings=self.type_vocab_size, + embedding_dim=self.hidden_size, + E_init=_truncated_normal(), + ) + self.embedding_norm = tlx.nn.LayerNorm( + normalized_shape=self.hidden_size, + epsilon=self.layer_norm_eps, + gamma_init="ones", + beta_init="zeros", + ) + self.embedding_dropout = tlx.nn.Dropout(p=self.dropout_rate) + self.encoder_layers = tlx.nn.ModuleList([ + GraphTokenTransformerLayer( + hidden_size=self.hidden_size, + num_attention_heads=self.num_attention_heads, + intermediate_size=self.intermediate_size, + dropout_rate=self.dropout_rate, + attention_dropout_rate=self.attention_dropout_rate, + layer_norm_eps=self.layer_norm_eps, + ) + for _ in range(self.num_hidden_layers) + ]) + # Retain the standard BERT pooler parameters used by the canonical architecture. + self.encoder_pooler = _linear(self.hidden_size, self.hidden_size) + self.mlm_head = _linear(self.hidden_size, self.vocab_size, bias=False) + task_hidden_size = max(1, self.hidden_size // 2) + self.task_head_in = _linear(self.hidden_size, task_hidden_size) + self.task_head_dropout = tlx.nn.Dropout(p=self.task_dropout_rate) + self.task_head_out = _linear(task_hidden_size, self.output_dim) + + def _embed(self, input_ids): + shape = tlx.get_tensor_shape(input_ids) + sequence_length = int(shape[1]) + if sequence_length > self.max_position_embeddings: + raise ValueError("Input sequence length exceeds max_position_embeddings.") + word_embeddings = self.token_embeddings(input_ids) + positions = tlx.cumsum(tlx.ones_like(input_ids), axis=1) - 1 + position_embeddings = self.position_embeddings(positions) + token_type_ids = input_ids * 0 + token_type_embeddings = self.token_type_embeddings(token_type_ids) + hidden_states = ( + word_embeddings + position_embeddings + token_type_embeddings) + return self.embedding_dropout(self.embedding_norm(hidden_states)) + + def _pool(self, hidden_states, attention_mask=None): + if self.pooling == "cls": + return hidden_states[:, 0] + return masked_mean(hidden_states, attention_mask) + + def _task_logits(self, pooled_output): + hidden = tlx.relu(self.task_head_in(pooled_output)) + hidden = self.task_head_dropout(hidden) + return self.task_head_out(hidden) + + def forward(self, input_ids, attention_mask=None, task: str | None = None): + if task not in {None, "mlm", "supervised"}: + raise ValueError("task must be None, 'mlm', or 'supervised'.") + hidden_states = self._embed(input_ids) + for layer in self.encoder_layers: + hidden_states = layer(hidden_states, attention_mask=attention_mask) + if task == "mlm": + return self.mlm_head(hidden_states) + pooled_output = self._pool(hidden_states, attention_mask) + if task == "supervised": + return self._task_logits(pooled_output) + mlm_logits = self.mlm_head(hidden_states) + task_logits = self._task_logits(pooled_output) + return { + "logits": task_logits, + "mlm_logits": mlm_logits, + "pooled_output": pooled_output, + "last_hidden_state": hidden_states, + } + + def manifest(self): + return { + "encoder_type": self.encoder_type, + "model_name": self.model_name, + "pooling": self.pooling, + "vocab_size": self.vocab_size, + "hidden_size": self.hidden_size, + "max_position_embeddings": self.max_position_embeddings, + "total_parameters": _parameter_count(self.all_weights), + "trainable_parameters": _parameter_count(self.trainable_weights), + "num_hidden_layers": self.num_hidden_layers, + "num_attention_heads": self.num_attention_heads, + "intermediate_size": self.intermediate_size, + "hidden_act": self.hidden_act, + "hidden_dropout_prob": self.dropout_rate, + "attention_probs_dropout_prob": self.attention_dropout_rate, + "position_embedding_type": self.position_embedding_type, + "layer_norm_eps": self.layer_norm_eps, + "framework": "tensorlayerx", + } + + +class GraphBERT(GraphTokenTransformer): + """Strict BERT-Small GraphTokenizer encoder.""" diff --git a/gammagl/models/graph_gte.py b/gammagl/models/graph_gte.py new file mode 100644 index 000000000..492574de7 --- /dev/null +++ b/gammagl/models/graph_gte.py @@ -0,0 +1,471 @@ +"""Native TensorLayerX GTE encoder for serialized graph tokens.""" + +from __future__ import annotations + +import math + +import tensorlayerx as tlx + +try: + from .graph_bert import ( + _linear, + _parameter_count, + _truncated_normal, + additive_attention_bias, + masked_mean, + ) +except ImportError: # Allows direct file loading in focused tests. + from graph_bert import ( + _linear, + _parameter_count, + _truncated_normal, + additive_attention_bias, + masked_mean, + ) + + +def rotate_half(value): + width = int(tlx.get_tensor_shape(value)[-1]) + if width % 2 != 0: + raise ValueError("RoPE head dimension must be even.") + half = width // 2 + first = value[..., :half] + second = value[..., half:] + return tlx.concat([-second, first], axis=-1) + + +def apply_rotary_pos_emb(query, key, cosine, sine): + return ( + query * cosine + rotate_half(query) * sine, + key * cosine + rotate_half(key) * sine, + ) + + +class RotaryEmbedding(tlx.nn.Module): + """Official-GTE-compatible NTK-scaled rotary position embedding.""" + + def __init__( + self, + dim: int, + max_position_embeddings: int = 8192, + base: float = 20000.0, + scaling_factor: float = 8.0, + name: str | None = None, + ): + super().__init__(name=name) + if dim <= 0 or dim % 2 != 0: + raise ValueError("RoPE dimension must be a positive even integer.") + if scaling_factor <= 0: + raise ValueError("RoPE scaling_factor must be positive.") + self.dim = int(dim) + self.max_position_embeddings = int(max_position_embeddings) + self.base = float(base) + self.scaling_factor = float(scaling_factor) + + def forward(self, reference): + sequence_length = int(tlx.get_tensor_shape(reference)[1]) + if sequence_length > self.max_position_embeddings: + raise ValueError("Input sequence length exceeds max_position_embeddings.") + half_reference = reference[0, 0, : self.dim // 2] + dimension_steps = ( + tlx.cumsum(tlx.ones_like(half_reference), axis=0) - 1.0) * 2.0 + exponent = dimension_steps / float(self.dim) + adjusted_base = self.base * self.scaling_factor + base_tensor = tlx.ones_like(dimension_steps) * adjusted_base + inverse_frequency = 1.0 / tlx.pow(base_tensor, exponent) + inverse_frequency = inverse_frequency / ( + self.scaling_factor ** (2.0 / float(self.dim))) + position_reference = reference[0, :, 0] + positions = tlx.cumsum(tlx.ones_like(position_reference), axis=0) - 1.0 + frequencies = tlx.einsum("i,j->ij", positions, inverse_frequency) + angles = tlx.concat([frequencies, frequencies], axis=-1) + cosine = tlx.expand_dims(tlx.expand_dims(tlx.cos(angles), axis=0), axis=0) + sine = tlx.expand_dims(tlx.expand_dims(tlx.sin(angles), axis=0), axis=0) + return cosine, sine + + +class GraphGTESelfAttention(tlx.nn.Module): + """Packed-QKV self-attention with RoPE applied to query and key.""" + + def __init__( + self, + hidden_size: int, + num_attention_heads: int, + attention_dropout_rate: float = 0.1, + name: str | None = None, + ): + super().__init__(name=name) + if hidden_size % num_attention_heads != 0: + raise ValueError("hidden_size must be divisible by num_attention_heads.") + self.hidden_size = int(hidden_size) + self.num_attention_heads = int(num_attention_heads) + self.head_size = self.hidden_size // self.num_attention_heads + if self.head_size % 2 != 0: + raise ValueError("GTE attention head size must be even for RoPE.") + self.qkv_proj = _linear(self.hidden_size, self.hidden_size * 3) + self.out_proj = _linear(self.hidden_size, self.hidden_size) + self.attention_dropout = tlx.nn.Dropout(p=float(attention_dropout_rate)) + + def _split_heads(self, value): + batch_size, sequence_length = tlx.get_tensor_shape(value)[:2] + value = tlx.reshape( + value, + [batch_size, sequence_length, self.num_attention_heads, self.head_size], + ) + return tlx.transpose(value, [0, 2, 1, 3]) + + def _merge_heads(self, value): + value = tlx.transpose(value, [0, 2, 1, 3]) + batch_size, sequence_length = tlx.get_tensor_shape(value)[:2] + return tlx.reshape(value, [batch_size, sequence_length, self.hidden_size]) + + def forward(self, hidden_states, attention_mask=None, rope_embeddings=None): + if rope_embeddings is None: + raise ValueError("GTE attention requires rotary embeddings.") + query, key, value = tlx.split( + self.qkv_proj(hidden_states), 3, axis=-1) + query = self._split_heads(query) + key = self._split_heads(key) + value = self._split_heads(value) + query, key = apply_rotary_pos_emb( + query, key, rope_embeddings[0], rope_embeddings[1]) + key_transposed = tlx.transpose(key, [0, 1, 3, 2]) + scores = tlx.matmul(query, key_transposed) / math.sqrt(self.head_size) + bias = additive_attention_bias(attention_mask) + if bias is not None: + scores = scores + bias + probabilities = self.attention_dropout(tlx.softmax(scores, axis=-1)) + context = tlx.matmul(probabilities, value) + return self.out_proj(self._merge_heads(context)) + + +class GraphGatedMLP(tlx.nn.Module): + """GTE gated GELU feed-forward network.""" + + def __init__( + self, + hidden_size: int, + intermediate_size: int, + dropout_rate: float = 0.1, + name: str | None = None, + ): + super().__init__(name=name) + self.intermediate_size = int(intermediate_size) + self.up_gate_proj = _linear( + hidden_size, self.intermediate_size * 2, bias=False) + self.down_proj = _linear(self.intermediate_size, hidden_size) + self.hidden_dropout = tlx.nn.Dropout(p=float(dropout_rate)) + + def forward(self, hidden_states): + up_states, gate = tlx.split( + self.up_gate_proj(hidden_states), 2, axis=-1) + gated_states = up_states * tlx.gelu(gate) + return self.down_proj(self.hidden_dropout(gated_states)) + + +class GraphGTELayer(tlx.nn.Module): + """Post-normalization GTE attention and gated-MLP block.""" + + def __init__( + self, + hidden_size: int, + num_attention_heads: int, + intermediate_size: int, + dropout_rate: float = 0.1, + attention_dropout_rate: float | None = None, + layer_norm_eps: float = 1e-12, + name: str | None = None, + ): + super().__init__(name=name) + attention_dropout_rate = ( + dropout_rate if attention_dropout_rate is None + else attention_dropout_rate + ) + self.attention = GraphGTESelfAttention( + hidden_size=hidden_size, + num_attention_heads=num_attention_heads, + attention_dropout_rate=attention_dropout_rate, + ) + self.mlp = GraphGatedMLP( + hidden_size=hidden_size, + intermediate_size=intermediate_size, + dropout_rate=dropout_rate, + ) + self.attention_norm = tlx.nn.LayerNorm( + normalized_shape=hidden_size, + epsilon=layer_norm_eps, + gamma_init="ones", + beta_init="zeros", + ) + self.mlp_norm = tlx.nn.LayerNorm( + normalized_shape=hidden_size, + epsilon=layer_norm_eps, + gamma_init="ones", + beta_init="zeros", + ) + self.hidden_dropout = tlx.nn.Dropout(p=float(dropout_rate)) + + def forward(self, hidden_states, attention_mask=None, rope_embeddings=None): + attention_output = self.attention( + hidden_states, + attention_mask=attention_mask, + rope_embeddings=rope_embeddings, + ) + hidden_states = self.attention_norm( + hidden_states + self.hidden_dropout(attention_output)) + mlp_output = self.mlp(hidden_states) + return self.mlp_norm(hidden_states + self.hidden_dropout(mlp_output)) + + +class GraphGTE(tlx.nn.Module): + """Native TLX GTE encoder; ``from_pretrained`` loads its GTE encoder.""" + + encoder_type = "gte" + model_name = "gte-base" + position_embedding_type = "rope" + hidden_act = "gelu" + + def __init__( + self, + vocab_size: int, + output_dim: int, + hidden_size: int = 768, + num_hidden_layers: int = 12, + num_attention_heads: int = 12, + intermediate_size: int = 3072, + max_position_embeddings: int = 8192, + dropout_rate: float = 0.1, + attention_dropout_rate: float | None = None, + layer_norm_eps: float = 1e-12, + pad_token_id: int = 0, + type_vocab_size: int = 1, + task_type: str = "regression", + pooling: str = "mean", + task_dropout: float = 0.1, + rope_theta: float = 20000.0, + rope_scaling_factor: float = 8.0, + name: str | None = None, + ): + super().__init__(name=name) + if pooling not in {"mean", "cls"}: + raise ValueError("pooling must be 'mean' or 'cls'.") + if hidden_size % num_attention_heads != 0: + raise ValueError("hidden_size must be divisible by num_attention_heads.") + self.vocab_size = int(vocab_size) + self.output_dim = int(output_dim) + self.hidden_size = int(hidden_size) + self.num_hidden_layers = int(num_hidden_layers) + self.num_attention_heads = int(num_attention_heads) + self.intermediate_size = int(intermediate_size) + self.max_position_embeddings = int(max_position_embeddings) + self.dropout_rate = float(dropout_rate) + self.attention_dropout_rate = float( + dropout_rate if attention_dropout_rate is None + else attention_dropout_rate) + self.layer_norm_eps = float(layer_norm_eps) + self.pad_token_id = int(pad_token_id) + self.type_vocab_size = int(type_vocab_size) + self.task_type = str(task_type) + self.pooling = pooling + self.task_dropout_rate = float(task_dropout) + self.rope_theta = float(rope_theta) + self.rope_scaling_factor = float(rope_scaling_factor) + self.head_size = self.hidden_size // self.num_attention_heads + + self.token_embeddings = tlx.nn.Embedding( + num_embeddings=self.vocab_size, + embedding_dim=self.hidden_size, + E_init=_truncated_normal(), + ) + self.token_type_embeddings = tlx.nn.Embedding( + num_embeddings=self.type_vocab_size, + embedding_dim=self.hidden_size, + E_init=_truncated_normal(), + ) + self.embedding_norm = tlx.nn.LayerNorm( + normalized_shape=self.hidden_size, + epsilon=self.layer_norm_eps, + gamma_init="ones", + beta_init="zeros", + ) + self.embedding_dropout = tlx.nn.Dropout(p=self.dropout_rate) + self.rotary_embedding = RotaryEmbedding( + dim=self.head_size, + max_position_embeddings=self.max_position_embeddings, + base=self.rope_theta, + scaling_factor=self.rope_scaling_factor, + ) + self.encoder_layers = tlx.nn.ModuleList([ + GraphGTELayer( + hidden_size=self.hidden_size, + num_attention_heads=self.num_attention_heads, + intermediate_size=self.intermediate_size, + dropout_rate=self.dropout_rate, + attention_dropout_rate=self.attention_dropout_rate, + layer_norm_eps=self.layer_norm_eps, + ) + for _ in range(self.num_hidden_layers) + ]) + self.mlm_head = _linear(self.hidden_size, self.vocab_size, bias=False) + task_hidden_size = max(1, self.hidden_size // 2) + self.task_head_in = _linear(self.hidden_size, task_hidden_size) + self.task_head_dropout = tlx.nn.Dropout(p=self.task_dropout_rate) + self.task_head_out = _linear(task_hidden_size, self.output_dim) + self.pretrained_manifest = None + + @classmethod + def from_pretrained(cls, vocab_size: int, output_dim: int, cache_dir=None, **kwargs): + """Build native TLX GraphGTE with official encoder weights only.""" + from .graph_gte_pretrained import ( + GTE_CHECKPOINT, GTE_MODEL_ID, GTE_REVISION, + download_official_gte_checkpoint, load_official_gte_config, + load_pretrained_encoder, + ) + config = load_official_gte_config(cache_dir=cache_dir) + architecture = { + "hidden_size": config["hidden_size"], + "num_hidden_layers": config["num_hidden_layers"], + "num_attention_heads": config["num_attention_heads"], + "intermediate_size": config["intermediate_size"], + "max_position_embeddings": config["max_position_embeddings"], + "dropout_rate": config["hidden_dropout_prob"], + "attention_dropout_rate": config["attention_probs_dropout_prob"], + "layer_norm_eps": config["layer_norm_eps"], + "pad_token_id": config["pad_token_id"], + "type_vocab_size": config["type_vocab_size"], + "rope_theta": config["rope_theta"], + "rope_scaling_factor": config["rope_scaling"]["factor"], + } + for key, value in kwargs.items(): + if key in architecture and architecture[key] != value: + raise RuntimeError( + f"Official GTE config mismatch for {key}: " + f"expected {architecture[key]!r}, got {value!r}.") + architecture.update(kwargs) + model = cls(vocab_size=vocab_size, output_dim=output_dim, **architecture) + path, actual_hash = download_official_gte_checkpoint(cache_dir=cache_dir) + report = load_pretrained_encoder(model, path) + model.pretrained_manifest = { + "model_id": GTE_MODEL_ID, "revision": GTE_REVISION, + "checkpoint_files": [GTE_CHECKPOINT], "checkpoint_sha256": {GTE_CHECKPOINT: actual_hash}, + "pretrained": True, + "encoder_expected_tensors": len(report["expected"]), + "encoder_loaded_tensors": len(report["loaded"]), + "encoder_coverage": report["coverage"], + "missing": report["missing"], "unexpected": report["unexpected"], + "ignored": report["ignored"], + "ignored_checkpoint_keys": report["ignored"], + "shape_mismatches": report["shape_mismatches"], + "graph_embedding_init": "new_truncated_normal(stddev=0.02)", + "task_head_init": "new_truncated_normal(stddev=0.02)", + "reproduction": True, + } + return model + + def _embed(self, input_ids): + sequence_length = int(tlx.get_tensor_shape(input_ids)[1]) + if sequence_length > self.max_position_embeddings: + raise ValueError("Input sequence length exceeds max_position_embeddings.") + word_embeddings = self.token_embeddings(input_ids) + token_type_embeddings = self.token_type_embeddings(input_ids * 0) + return self.embedding_dropout( + self.embedding_norm(word_embeddings + token_type_embeddings)) + + def _embed_inputs(self, inputs_embeds): + """Apply the official GTE embedding post-processing to raw embeddings. + + This intentionally accepts raw embeddings rather than graph token IDs, + so the official Hugging Face encoder can be used only as a test oracle. + """ + sequence_length = int(tlx.get_tensor_shape(inputs_embeds)[1]) + if sequence_length > self.max_position_embeddings: + raise ValueError("Input sequence length exceeds max_position_embeddings.") + token_type_ids = tlx.zeros( + [int(tlx.get_tensor_shape(inputs_embeds)[0]), sequence_length], + dtype=tlx.int64, + ) + return self.embedding_dropout(self.embedding_norm( + inputs_embeds + self.token_type_embeddings(token_type_ids))) + + def encode_embeddings(self, inputs_embeds, attention_mask=None, + return_hidden_states=False): + """Run the native encoder from raw embeddings for reference testing.""" + hidden_states = self._embed_inputs(inputs_embeds) + hidden_states_list = [hidden_states] + rope_embeddings = self.rotary_embedding(hidden_states) + for layer in self.encoder_layers: + hidden_states = layer(hidden_states, attention_mask, rope_embeddings) + hidden_states_list.append(hidden_states) + if return_hidden_states: + return hidden_states, tuple(hidden_states_list) + return hidden_states + + def _pool(self, hidden_states, attention_mask=None): + if self.pooling == "cls": + return hidden_states[:, 0] + return masked_mean(hidden_states, attention_mask) + + def _task_logits(self, pooled_output): + hidden = tlx.relu(self.task_head_in(pooled_output)) + hidden = self.task_head_dropout(hidden) + return self.task_head_out(hidden) + + def forward(self, input_ids, attention_mask=None, task: str | None = None): + if task not in {None, "mlm", "supervised"}: + raise ValueError("task must be None, 'mlm', or 'supervised'.") + hidden_states = self._embed(input_ids) + rope_embeddings = self.rotary_embedding(hidden_states) + for layer in self.encoder_layers: + hidden_states = layer(hidden_states, attention_mask, rope_embeddings) + if task == "mlm": + return self.mlm_head(hidden_states) + pooled_output = self._pool(hidden_states, attention_mask) + if task == "supervised": + return self._task_logits(pooled_output) + mlm_logits = self.mlm_head(hidden_states) + task_logits = self._task_logits(pooled_output) + return { + "logits": task_logits, + "mlm_logits": mlm_logits, + "pooled_output": pooled_output, + "last_hidden_state": hidden_states, + } + + def manifest(self): + manifest = { + "encoder_type": self.encoder_type, + "model_name": self.model_name, + "pooling": self.pooling, + "vocab_size": self.vocab_size, + "hidden_size": self.hidden_size, + "max_position_embeddings": self.max_position_embeddings, + "total_parameters": _parameter_count(self.all_weights), + "trainable_parameters": _parameter_count(self.trainable_weights), + "num_hidden_layers": self.num_hidden_layers, + "num_attention_heads": self.num_attention_heads, + "intermediate_size": self.intermediate_size, + "hidden_act": self.hidden_act, + "hidden_dropout_prob": self.dropout_rate, + "attention_probs_dropout_prob": self.attention_dropout_rate, + "position_embedding_type": self.position_embedding_type, + "layer_norm_eps": self.layer_norm_eps, + "rope_theta": self.rope_theta, + "rope_scaling": { + "type": "ntk", + "factor": self.rope_scaling_factor, + }, + "framework": "tensorlayerx", + "weight_source": "random_initialization", + "official_checkpoint": None, + } + if self.pretrained_manifest is not None: + manifest.update(self.pretrained_manifest) + manifest["weight_source"] = "official_gte_encoder" + else: + manifest.update({ + "pretrained": False, + "reproduction": False, + "graph_embedding_init": "new_truncated_normal(stddev=0.02)", + "task_head_init": "new_truncated_normal(stddev=0.02)", + }) + return manifest diff --git a/gammagl/models/graph_gte_pretrained.py b/gammagl/models/graph_gte_pretrained.py new file mode 100644 index 000000000..2b6513723 --- /dev/null +++ b/gammagl/models/graph_gte_pretrained.py @@ -0,0 +1,129 @@ +"""Strict official GTE checkpoint loading for the native TLX GraphGTE.""" +import hashlib +import json +from pathlib import Path + + +GTE_MODEL_ID = "Alibaba-NLP/gte-multilingual-base" +GTE_REVISION = "9bbca17d9273fd0d03d5725c7a4b0f6b45142062" +GTE_CHECKPOINT = "model.safetensors" +GTE_CHECKPOINT_SHA256 = "f5a35a10faa54da7717870af1517c9b41e9bd8e3880bc5a8e9363d4c3c63e9b0" + + +def sha256_file(path, chunk_size=1024 * 1024): + digest = hashlib.sha256() + with Path(path).open("rb") as handle: + for chunk in iter(lambda: handle.read(chunk_size), b""): + digest.update(chunk) + return digest.hexdigest() + + +def download_official_gte_checkpoint(cache_dir=None): + from huggingface_hub import hf_hub_download + path = hf_hub_download(GTE_MODEL_ID, GTE_CHECKPOINT, revision=GTE_REVISION, + cache_dir=cache_dir) + actual = sha256_file(path) + if actual != GTE_CHECKPOINT_SHA256: + raise RuntimeError( + f"GTE checkpoint SHA-256 mismatch: expected {GTE_CHECKPOINT_SHA256}, got {actual}.") + return Path(path), actual + + +def load_official_gte_config(cache_dir=None): + from huggingface_hub import hf_hub_download + path = hf_hub_download(GTE_MODEL_ID, "config.json", revision=GTE_REVISION, + cache_dir=cache_dir) + with open(path, encoding="utf-8") as handle: + config = json.load(handle) + required = {"hidden_size": 768, "num_hidden_layers": 12, + "num_attention_heads": 12, "intermediate_size": 3072, + "pack_qkv": True, "position_embedding_type": "rope", + "rope_theta": 20000, "max_position_embeddings": 8192, + "hidden_act": "gelu", "layer_norm_type": "layer_norm", + "attention_probs_dropout_prob": 0.0, + "use_memory_efficient_attention": False} + for key, value in required.items(): + if config.get(key) != value: + raise RuntimeError(f"Official GTE config mismatch for {key}: {config.get(key)!r}.") + if config.get("rope_scaling") != {"factor": 8.0, "type": "ntk"}: + raise RuntimeError("Official GTE RoPE scaling config mismatch.") + return config + + +def _assign_tensor(target, value, key, transpose=False): + """Assign one semantically mapped PyTorch tensor to its TLX variable.""" + import tensorlayerx as tlx + + value = value.detach().cpu().numpy() + if transpose: + if value.ndim != 2: + raise ValueError(f"Only matrix mappings may transpose: {key}.") + value = value.T + if tuple(value.shape) != tuple(target.shape): + raise ValueError( + f"{key}: mapped shape {tuple(value.shape)} does not match " + f"TLX target {tuple(target.shape)}.") + value = tlx.convert_to_tensor(value, dtype=target.dtype) + if hasattr(target, "assign"): + target.assign(value) + else: + target.data.copy_(value) + + +def load_pretrained_encoder(model, checkpoint_path): + """Load the explicitly mapped official encoder tensors into ``GraphGTE``.""" + from safetensors import safe_open + + with safe_open(str(checkpoint_path), framework="pt", device="cpu") as handle: + source = {key: handle.get_tensor(key) for key in handle.keys()} + ignored = { + "new.embeddings.word_embeddings.weight": + "text-vocabulary embedding is incompatible with Graph BPE tokens", + "classifier.weight": "official classification head is not a GraphTokenizer task head", + "classifier.bias": "official classification head is not a GraphTokenizer task head", + } + # Every entry is a source-key-to-TLX-variable semantic mapping. The + # official NewModel and TLX torch backend both store Linear kernels as + # [out_features, in_features], hence transpose=False is deliberate rather + # than inferred from shape. + mapping = { + "new.embeddings.token_type_embeddings.weight": (model.token_type_embeddings.embeddings, False, "embedding"), + "new.embeddings.LayerNorm.weight": (model.embedding_norm.gamma, False, "layer_norm_gamma"), + "new.embeddings.LayerNorm.bias": (model.embedding_norm.beta, False, "layer_norm_beta"), + } + for index, layer in enumerate(model.encoder_layers): + prefix = f"new.encoder.layer.{index}" + mapping.update({ + f"{prefix}.attention.qkv_proj.weight": (layer.attention.qkv_proj.weights, False, "packed_qkv_QKV"), + f"{prefix}.attention.qkv_proj.bias": (layer.attention.qkv_proj.biases, False, "packed_qkv_bias_QKV"), + f"{prefix}.attention.o_proj.weight": (layer.attention.out_proj.weights, False, "attention_output"), + f"{prefix}.attention.o_proj.bias": (layer.attention.out_proj.biases, False, "attention_output_bias"), + f"{prefix}.attn_ln.weight": (layer.attention_norm.gamma, False, "post_attention_layer_norm_gamma"), + f"{prefix}.attn_ln.bias": (layer.attention_norm.beta, False, "post_attention_layer_norm_beta"), + f"{prefix}.mlp.up_gate_proj.weight": (layer.mlp.up_gate_proj.weights, False, "packed_up_gate_up_then_gate"), + f"{prefix}.mlp.down_proj.weight": (layer.mlp.down_proj.weights, False, "mlp_down"), + f"{prefix}.mlp.down_proj.bias": (layer.mlp.down_proj.biases, False, "mlp_down_bias"), + f"{prefix}.mlp_ln.weight": (layer.mlp_norm.gamma, False, "post_mlp_layer_norm_gamma"), + f"{prefix}.mlp_ln.bias": (layer.mlp_norm.beta, False, "post_mlp_layer_norm_beta"), + }) + expected = set(mapping) + missing = sorted(expected - set(source)) + unexpected = sorted(set(source) - expected - set(ignored)) + mismatches = [] + loaded = [] + for key, (target, transpose, _semantic) in mapping.items(): + if key not in source: + continue + try: + _assign_tensor(target, source[key], key, transpose=transpose) + except ValueError: + mismatches.append({"key": key, "source": tuple(source[key].shape), + "target": tuple(target.shape), "transpose": transpose}) + continue + loaded.append(key) + report = {"expected": sorted(expected), "loaded": sorted(loaded), "missing": missing, + "unexpected": unexpected, "ignored": ignored, "shape_mismatches": mismatches, + "coverage": len(loaded) / len(expected) if expected else 1.0} + if missing or unexpected or mismatches or report["coverage"] != 1.0: + raise RuntimeError(f"Incomplete GTE encoder conversion: {report}") + return report diff --git a/gammagl/transforms/__init__.py b/gammagl/transforms/__init__.py index b7ff5dfbd..40c06a29a 100644 --- a/gammagl/transforms/__init__.py +++ b/gammagl/transforms/__init__.py @@ -7,6 +7,9 @@ from .random_link_split import RandomLinkSplit from .vgae_pre import mask_test_edges, sparse_to_tuple from .svd_feature_reduction import SVDFeatureReduction +from .graph_serializer import EulerianSerializer, FrequencyGuidedEulerianSerializer, GraphSerializer +from .graph_bpe import BPECodebook, GraphBPE, BPEEngine +from .graph_tokenizer import GraphTokenizer, GraphTokenizationResult, GraphTokenizerMLMBatch, GraphTokenizerSpecialTokens __all__ = [ 'BaseTransform', @@ -18,7 +21,17 @@ 'RandomLinkSplit', 'mask_test_edges', 'sparse_to_tuple', - 'SVDFeatureReduction' + 'SVDFeatureReduction', + 'FrequencyGuidedEulerianSerializer', + 'EulerianSerializer', + 'GraphSerializer', + 'BPECodebook', + 'GraphBPE', + 'BPEEngine', + 'GraphTokenizer', + 'GraphTokenizationResult', + 'GraphTokenizerMLMBatch', + 'GraphTokenizerSpecialTokens' ] diff --git a/gammagl/transforms/graph_bpe.py b/gammagl/transforms/graph_bpe.py new file mode 100644 index 000000000..7cd88f13f --- /dev/null +++ b/gammagl/transforms/graph_bpe.py @@ -0,0 +1,276 @@ +import json +import pickle +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Dict, Iterable, List, Tuple + + +MergeRule = Tuple[int, int, int] + + +@dataclass +class BPECodebook: + """Serializable BPE merge table.""" + + merge_rules: List[MergeRule] = field(default_factory=list) + vocab_size: int = 0 + metadata: Dict[str, Any] = field(default_factory=dict) + + +class GraphBPE: + """Minimal public interface for graph-token BPE training and encoding.""" + + def __init__( + self, + num_merges: int = 2000, + min_frequency: int = 2, + backend: str = "python", + protected_token_ids=None, + ): + self.num_merges = int(num_merges) + self.min_frequency = int(min_frequency) + self.backend = backend + self.protected_token_ids = { + int(token) for token in (protected_token_ids or ())} + self.codebook = BPECodebook() + + def fit(self, token_sequences): + sequences = [self._as_int_list(sequence) for sequence in token_sequences] + training_sequences = self._training_segments(sequences) + if self.backend == "cpp": + return self._fit_with_bridge(training_sequences, sequences) + if self.backend == "auto": + from third_party import graph_bpe_cpp + if graph_bpe_cpp.is_available(): + return self._fit_with_bridge(training_sequences, sequences) + elif self.backend != "python": + raise ValueError("backend must be one of: python, auto, cpp.") + + return self._fit_python(training_sequences, sequences) + + def _fit_python(self, token_sequences, original_sequences=None): + sequences = [self._as_int_list(sequence) for sequence in token_sequences] + merge_rules: List[MergeRule] = [] + next_token_id = self._next_token_id(original_sequences or sequences) + + for _ in range(self.num_merges): + pair_counts = self._count_pairs(sequences) + if not pair_counts: + break + + best_pair, best_count = self._select_best_pair(pair_counts) + if best_count < self.min_frequency: + break + + merge_rule = (best_pair[0], best_pair[1], next_token_id) + merge_rules.append(merge_rule) + sequences = [self._apply_merge(sequence, merge_rule) for sequence in sequences] + next_token_id += 1 + + self.codebook = BPECodebook( + merge_rules=merge_rules, + vocab_size=next_token_id, + metadata={ + "num_merges_requested": self.num_merges, + "num_merges_performed": len(merge_rules), + "min_frequency": self.min_frequency, + "backend": "python", + }, + ) + return self + + def encode(self, token_sequence): + bridge = self._native_bridge() + pieces = self._split_on_protected(self._as_int_list(token_sequence)) + encoded = [] + for protected, piece in pieces: + if protected: + encoded.extend(piece) + elif bridge is not None and self.codebook.merge_rules: + encoded.extend(bridge.encode(piece, self.codebook.merge_rules)) + else: + for merge_rule in self.codebook.merge_rules: + piece = self._apply_merge(piece, merge_rule) + encoded.extend(piece) + return encoded + + def batch_encode(self, token_sequences): + bridge = self._native_bridge() + if not self.protected_token_ids and bridge is not None and self.codebook.merge_rules: + return bridge.batch_encode(token_sequences, self.codebook.merge_rules) + return [self.encode(sequence) for sequence in token_sequences] + + def decode(self, token_sequence): + """Expand BPE merge tokens back to the original token sequence.""" + expansions = { + int(merged): (int(left), int(right)) + for left, right, merged in self.codebook.merge_rules + } + decoded = [] + for token in self._as_int_list(token_sequence): + decoded.extend(self._expand_token(token, expansions, set())) + return decoded + + @classmethod + def _expand_token(cls, token: int, expansions, visiting) -> List[int]: + if token not in expansions: + return [token] + if token in visiting: + raise ValueError(f"BPE merge rules contain a cycle at token {token}.") + left, right = expansions[token] + visiting.add(token) + expanded = cls._expand_token(left, expansions, visiting) + cls._expand_token( + right, expansions, visiting) + visiting.remove(token) + return expanded + + def _native_bridge(self): + if self.backend != "cpp" and self.codebook.metadata.get("backend") != "cpp": + return None + from third_party import graph_bpe_cpp + if graph_bpe_cpp.is_available(): + return graph_bpe_cpp + if self.backend == "cpp": + raise ImportError( + "GraphBPE backend='cpp' requires the optional graph_bpe_cpp native extension. " + "Use backend='auto' or backend='python' to allow fallback.") + return None + + def save_codebook(self, path) -> None: + path = Path(path) + data = { + "merge_rules": [list(rule) for rule in self.codebook.merge_rules], + "vocab_size": self.codebook.vocab_size, + "metadata": dict(self.codebook.metadata), + } + if path.suffix.lower() in {".pkl", ".pickle"}: + with path.open("wb") as handle: + pickle.dump(data, handle) + return + if path.suffix.lower() == ".json": + with path.open("w", encoding="utf-8") as handle: + json.dump(data, handle, ensure_ascii=False) + return + raise ValueError("Codebook path must end with .json, .pkl, or .pickle.") + + @classmethod + def load_codebook(cls, path, backend: str = "python") -> "GraphBPE": + path = Path(path) + if path.suffix.lower() in {".pkl", ".pickle"}: + with path.open("rb") as handle: + data = pickle.load(handle) + elif path.suffix.lower() == ".json": + with path.open("r", encoding="utf-8") as handle: + data = json.load(handle) + else: + raise ValueError("Codebook path must end with .json, .pkl, or .pickle.") + + engine = cls(num_merges=0, backend=backend) + engine.codebook = BPECodebook( + merge_rules=[tuple(int(value) for value in rule) for rule in data.get("merge_rules", [])], + vocab_size=int(data.get("vocab_size", 0)), + metadata=dict(data.get("metadata", {})), + ) + return engine + + def _fit_with_bridge(self, token_sequences, original_sequences=None): + from third_party import graph_bpe_cpp + + if self.backend == "cpp" and not graph_bpe_cpp.is_available(): + raise ImportError( + "GraphBPE backend='cpp' requires the optional graph_bpe_cpp native extension. " + "Use backend='auto' or backend='python' to allow fallback." + ) + + initial_vocab_size = self._next_token_id(original_sequences or token_sequences) + result = graph_bpe_cpp.train_bpe( + token_sequences, + num_merges=self.num_merges, + min_frequency=self.min_frequency, + initial_vocab_size=initial_vocab_size, + ) + self.codebook = BPECodebook( + merge_rules=[tuple(int(value) for value in rule) for rule in result["merge_rules"]], + vocab_size=int(result["vocab_size"]), + metadata=dict(result.get("metadata", {})), + ) + if "backend" not in self.codebook.metadata: + self.codebook.metadata["backend"] = graph_bpe_cpp.backend_name() + return self + + def _training_segments(self, sequences): + if not self.protected_token_ids: + return sequences + return [ + piece + for sequence in sequences + for protected, piece in self._split_on_protected(sequence) + if not protected and piece + ] + + def _split_on_protected(self, sequence): + if not self.protected_token_ids: + return [(False, sequence)] + pieces = [] + current = [] + for token in sequence: + if token in self.protected_token_ids: + if current: + pieces.append((False, current)) + current = [] + pieces.append((True, [token])) + else: + current.append(token) + if current: + pieces.append((False, current)) + return pieces + + @staticmethod + def _as_int_list(sequence) -> List[int]: + if hasattr(sequence, "tolist"): + sequence = sequence.tolist() + return [int(token) for token in sequence] + + @staticmethod + def _next_token_id(sequences: Iterable[List[int]]) -> int: + max_token = -1 + for sequence in sequences: + if sequence: + max_token = max(max_token, max(sequence)) + return max_token + 1 + + @staticmethod + def _count_pairs( + sequences: Iterable[List[int]], protected_token_ids=None, + ) -> Dict[Tuple[int, int], int]: + counts: Dict[Tuple[int, int], int] = {} + protected_token_ids = set(protected_token_ids or ()) + for sequence in sequences: + for left, right in zip(sequence, sequence[1:]): + if left in protected_token_ids or right in protected_token_ids: + continue + pair = (left, right) + counts[pair] = counts.get(pair, 0) + 1 + return counts + + @staticmethod + def _select_best_pair(pair_counts: Dict[Tuple[int, int], int]) -> Tuple[Tuple[int, int], int]: + best_pair = min(pair_counts.keys(), key=lambda pair: (-pair_counts[pair], pair[0], pair[1])) + return best_pair, pair_counts[best_pair] + + @staticmethod + def _apply_merge(sequence: List[int], merge_rule: MergeRule) -> List[int]: + left, right, new_id = merge_rule + merged: List[int] = [] + index = 0 + while index < len(sequence): + if index + 1 < len(sequence) and sequence[index] == left and sequence[index + 1] == right: + merged.append(new_id) + index += 2 + else: + merged.append(sequence[index]) + index += 1 + return merged + + +BPEEngine = GraphBPE diff --git a/gammagl/transforms/graph_serializer.py b/gammagl/transforms/graph_serializer.py new file mode 100644 index 000000000..64aa585d6 --- /dev/null +++ b/gammagl/transforms/graph_serializer.py @@ -0,0 +1,568 @@ +from collections import defaultdict, deque +from dataclasses import dataclass, field +from typing import Any, Dict, Iterable, List, Sequence, Tuple + + +@dataclass +class GraphSerializationResult: + """Container for serialized graph tokens and lightweight metadata.""" + + token_ids: List[int] + metadata: Dict[str, Any] = field(default_factory=dict) + + +class FrequencyGuidedEulerianSerializer: + """Reversible Feuler serializer for undirected simple graphs only. + + ``edge_index`` may contain one or both directions of an undirected edge, + but every direction must carry the same edge label. Directed graphs, + parallel edges, and self-loops are not representable by this serializer. + + The token grammar is self-describing: a component starts with a node-label + token and a node-reference token, then contains repeated + ``edge-label, node-label, node-reference`` triples. Components are joined + by ``component_sep_token_id``. Node labels, edge labels, and node + references occupy disjoint integer domains, so the original graph can be + recovered from the token stream alone. + """ + + def __init__( + self, + name: str = "feuler", + include_edge_tokens: bool = True, + component_sep_token_id: int = -1, + ): + self.name = name + self.include_edge_tokens = include_edge_tokens + self.component_sep_token_id = int(component_sep_token_id) + self.frequency_map: Dict[Any, int] = {} + + def fit(self, graphs: Sequence[Any]): + counts: Dict[Tuple[int, int, int], int] = defaultdict(int) + for graph in graphs: + graph_view = self._read_graph(graph) + for src, dst, edge_label in graph_view["canonical_edges"]: + counts[self._edge_pattern(graph_view, src, dst, edge_label)] += 1 + self.frequency_map = dict(counts) + return self + + def serialize(self, graph: Any) -> GraphSerializationResult: + graph_view = self._read_graph(graph) + components = self._connected_components(graph_view["num_nodes"], graph_view["arcs"]) + component_results = [] + + for component in components: + edges = self._serialize_component(graph_view, component) + tokens = self._edges_to_tokens(graph_view, edges, component) + component_results.append(( + tokens, + len(edges), + self._traversal_signature(graph_view, edges, component), + )) + + component_results.sort(key=lambda item: (-len(item[0]), item[2])) + token_ids: List[int] = [] + for index, (component_tokens, _, _) in enumerate(component_results): + if index > 0: + token_ids.append(self.component_sep_token_id) + token_ids.extend(component_tokens) + + return GraphSerializationResult( + token_ids=token_ids, + metadata={ + "method": self.name, + "protocol_version": 2, + "num_nodes": graph_view["num_nodes"], + "num_components": len(component_results), + "num_edges_traversed": sum( + edge_count for _, edge_count, _ in component_results), + }, + ) + + def restore_input_metadata(self, result: GraphSerializationResult) -> Dict[str, List[int]]: + """Backward-compatible alias for token-stream reconstruction.""" + return self.deserialize(result) + + def deserialize(self, result: GraphSerializationResult) -> Dict[str, List[int]]: + """Reconstruct a graph from a Feuler token stream without metadata.""" + if not self.include_edge_tokens: + raise ValueError("Feuler deserialization requires include_edge_tokens=True.") + token_ids = result.token_ids if isinstance(result, GraphSerializationResult) else result + components = self._split_components(token_ids) + labels_by_node: Dict[int, int] = {} + edges_by_key: Dict[Tuple[int, int], int] = {} + + for component in components: + if not component: + raise ValueError("Feuler token stream contains an empty component.") + if len(component) < 2: + raise ValueError("Feuler component must begin with a node label and reference.") + current, label = self._parse_node(component, 0) + self._record_node_label(labels_by_node, current, label) + index = 2 + while index < len(component): + if index + 2 >= len(component): + raise ValueError("Feuler edge traversal must contain edge and node tokens.") + edge_label = self._decode_edge_token(component[index]) + nxt, label = self._parse_node(component, index + 1) + self._record_node_label(labels_by_node, nxt, label) + if current == nxt: + raise ValueError("Feuler token stream contains an unsupported self-loop.") + key = (min(current, nxt), max(current, nxt)) + previous = edges_by_key.setdefault(key, edge_label) + if previous != edge_label: + raise ValueError( + f"Feuler token stream assigns conflicting labels to edge {key}.") + current = nxt + index += 3 + + if not labels_by_node: + return {"edge_index": [[], []], "x": [], "edge_attr": [], "num_nodes": 0} + num_nodes = max(labels_by_node) + 1 + expected_nodes = set(range(num_nodes)) + if set(labels_by_node) != expected_nodes: + raise ValueError("Feuler node references must be contiguous from zero.") + canonical_edges = [(src, dst, label) for (src, dst), label in sorted(edges_by_key.items())] + return { + "edge_index": [ + [src for src, _, _ in canonical_edges], + [dst for _, dst, _ in canonical_edges], + ], + "x": [labels_by_node[node] for node in range(num_nodes)], + "edge_attr": [label for _, _, label in canonical_edges], + "num_nodes": num_nodes, + } + + @staticmethod + def node_token(label: int) -> int: + return 3 * FrequencyGuidedEulerianSerializer._encode_integer(label) + + @staticmethod + def edge_token(label: int) -> int: + return 3 * FrequencyGuidedEulerianSerializer._encode_integer(label) + 1 + + @staticmethod + def node_reference_token(node_id: int) -> int: + if int(node_id) < 0: + raise ValueError("Feuler node references must be non-negative.") + return 3 * int(node_id) + 2 + + @staticmethod + def _encode_integer(value: int) -> int: + value = int(value) + return 2 * value if value >= 0 else -2 * value - 1 + + @staticmethod + def _decode_integer(value: int) -> int: + return value // 2 if value % 2 == 0 else -(value // 2) - 1 + + def _parse_node(self, tokens: Sequence[int], index: int) -> Tuple[int, int]: + return self._decode_node_reference_token(tokens[index + 1]), self._decode_node_token(tokens[index]) + + def _decode_node_token(self, token: int) -> int: + token = int(token) + if token < 0 or token % 3 != 0: + raise ValueError(f"Expected a Feuler node-label token, got {token}.") + return self._decode_integer(token // 3) + + def _decode_edge_token(self, token: int) -> int: + token = int(token) + if token < 0 or token % 3 != 1: + raise ValueError(f"Expected a Feuler edge-label token, got {token}.") + return self._decode_integer(token // 3) + + @staticmethod + def _decode_node_reference_token(token: int) -> int: + token = int(token) + if token < 0 or token % 3 != 2: + raise ValueError(f"Expected a Feuler node-reference token, got {token}.") + return token // 3 + + def _split_components(self, token_ids: Sequence[int]) -> List[List[int]]: + components = [[]] + for token in token_ids: + token = int(token) + if token == self.component_sep_token_id: + components.append([]) + else: + components[-1].append(token) + return [] if components == [[]] else components + + @staticmethod + def _record_node_label(labels_by_node: Dict[int, int], node: int, label: int) -> None: + previous = labels_by_node.setdefault(node, label) + if previous != label: + raise ValueError( + f"Feuler token stream assigns conflicting labels to node {node}.") + + def _read_graph(self, graph: Any) -> Dict[str, Any]: + self._validate_graph_kind(graph) + edge_index = self._get_graph_attr(graph, "edge_index") + if edge_index is None: + raise ValueError("Graph must provide edge_index.") + + edge_pairs = self._edge_pairs(edge_index) + num_nodes = self._infer_num_nodes(graph, edge_pairs) + for src, dst in edge_pairs: + if src < 0 or dst < 0 or src >= num_nodes or dst >= num_nodes: + raise ValueError( + "edge_index contains an endpoint outside the declared graph nodes.") + node_labels = self._node_labels(self._get_graph_attr(graph, "x"), num_nodes) + edge_labels = self._edge_labels(self._get_graph_attr(graph, "edge_attr"), len(edge_pairs)) + input_edges = [(int(src), int(dst), int(edge_labels[i])) for i, (src, dst) in enumerate(edge_pairs)] + canonical_edges = self._canonicalize_undirected_edges(input_edges) + arcs = self._undirected_arcs(canonical_edges) + return { + "num_nodes": num_nodes, + "node_labels": node_labels, + "canonical_edges": canonical_edges, + "arcs": arcs, + "node_signatures": self._node_structure_signatures( + num_nodes, node_labels, arcs), + } + + def _validate_graph_kind(self, graph: Any) -> None: + """Reject graph objects that explicitly declare directed semantics. + + A single COO orientation remains a supported *storage* form for one + undirected edge. It is indistinguishable from a directed edge without + graph metadata, so objects that expose ``directed`` or ``is_directed`` + must declare it false. + """ + for name in ("directed", "is_directed"): + value = self._get_graph_attr(graph, name) + if callable(value): + value = value() + if value is not None and bool(value): + raise ValueError( + "Feuler only supports undirected simple graphs; " + "directed graph semantics are unsupported.") + + @staticmethod + def _canonicalize_undirected_edges( + input_edges: Sequence[Tuple[int, int, int]], + ) -> List[Tuple[int, int, int]]: + """Validate COO input and normalize it to one oriented edge per pair.""" + labels_by_edge: Dict[Tuple[int, int], int] = {} + directed_edges = set() + for src, dst, label in input_edges: + if src == dst: + raise ValueError("Feuler does not support self-loops.") + directed_key = (src, dst) + if directed_key in directed_edges: + raise ValueError( + "Parallel edges are unsupported by the undirected simple " + "Feuler serializer.") + directed_edges.add(directed_key) + key = (min(src, dst), max(src, dst)) + previous = labels_by_edge.setdefault(key, label) + if previous != label: + raise ValueError( + "Feuler only supports undirected simple graphs; " + f"edge {key} has conflicting labels {previous} and {label}.") + return [(src, dst, label) for (src, dst), label in sorted(labels_by_edge.items())] + + @staticmethod + def _get_graph_attr(graph: Any, name: str): + if isinstance(graph, dict): + return graph.get(name) + return getattr(graph, name, None) + + def _infer_num_nodes(self, graph: Any, edge_pairs: Sequence[Tuple[int, int]]) -> int: + value = self._get_graph_attr(graph, "num_nodes") + if callable(value): + value = value() + if value is not None: + return int(value) + x = self._get_graph_attr(graph, "x") + if x is not None: + return len(self._to_list(x)) + if not edge_pairs: + return 0 + return max(max(src, dst) for src, dst in edge_pairs) + 1 + + def _edge_pairs(self, edge_index: Any) -> List[Tuple[int, int]]: + values = self._to_list(edge_index) + if len(values) == 2 and self._is_sequence(values[0]) and self._is_sequence(values[1]): + if len(values[0]) != len(values[1]): + raise ValueError("edge_index source and destination rows must have the same length.") + return [(int(src), int(dst)) for src, dst in zip(values[0], values[1])] + return [(int(src), int(dst)) for src, dst in values] + + def _node_labels(self, x: Any, num_nodes: int) -> List[int]: + if x is None: + return list(range(num_nodes)) + values = self._to_list(x) + if len(values) != num_nodes: + raise ValueError( + f"x must contain exactly {num_nodes} node features, got {len(values)}.") + return [self._scalar(value) for value in values] + + def _edge_labels(self, edge_attr: Any, num_edges: int) -> List[int]: + if edge_attr is None: + return [0] * num_edges + values = self._to_list(edge_attr) + if len(values) != num_edges: + raise ValueError( + f"edge_attr must contain exactly {num_edges} labels matching edge_index, got {len(values)}.") + return [self._scalar(value) for value in values] + + def _undirected_arcs(self, input_edges: Sequence[Tuple[int, int, int]]) -> List[Tuple[int, int, int]]: + arcs: List[Tuple[int, int, int]] = [] + for src, dst, label in input_edges: + arcs.append((src, dst, label)) + arcs.append((dst, src, label)) + return arcs + + def _connected_components(self, num_nodes: int, arcs: Sequence[Tuple[int, int, int]]) -> List[List[int]]: + neighbors: Dict[int, List[int]] = defaultdict(list) + for src, dst, _ in arcs: + neighbors[src].append(dst) + neighbors[dst].append(src) + + components: List[List[int]] = [] + visited = set() + for node in range(num_nodes): + if node in visited: + continue + queue = deque([node]) + visited.add(node) + component = [] + while queue: + current = queue.popleft() + component.append(current) + for nxt in neighbors[current]: + if nxt not in visited: + visited.add(nxt) + queue.append(nxt) + components.append(component) + return components + + def _serialize_component(self, graph_view: Dict[str, Any], component: Sequence[int]) -> List[Tuple[int, int, int]]: + component_nodes = set(component) + arcs = [ + (src, dst, label) + for src, dst, label in graph_view["arcs"] + if src in component_nodes and dst in component_nodes + ] + if not arcs: + return [] + + adjacency: Dict[int, List[Tuple[int, int, int, int]]] = defaultdict(list) + for arc_id, (src, dst, label) in enumerate(arcs): + adjacency[src].append((dst, label, arc_id, self._arc_priority(graph_view, src, dst, label))) + for src in adjacency: + adjacency[src].sort(key=lambda item: ( + -item[3], item[1], graph_view["node_signatures"][item[0]])) + adjacency[src].reverse() + + start = self._select_structural_node( + graph_view, (node for node in component if adjacency.get(node))) + used = set() + circuit: List[Tuple[int, int, int]] = [] + node_stack = [start] + edge_stack: List[Tuple[int, int, int]] = [] + while node_stack: + node = node_stack[-1] + while adjacency[node] and adjacency[node][-1][2] in used: + adjacency[node].pop() + if adjacency[node]: + dst, label, arc_id, _ = adjacency[node].pop() + used.add(arc_id) + node_stack.append(dst) + edge_stack.append((node, dst, label)) + continue + node_stack.pop() + if edge_stack: + circuit.append(edge_stack.pop()) + + circuit.reverse() + return circuit + + def _arc_priority(self, graph_view: Dict[str, Any], src: int, dst: int, edge_label: int) -> int: + pattern = self._edge_pattern(graph_view, src, dst, edge_label) + return int(self.frequency_map.get(pattern, 0)) + + @staticmethod + def _edge_pattern( + graph_view: Dict[str, Any], src: int, dst: int, edge_label: int, + ) -> Tuple[int, int, int]: + """Undirected frequency key; endpoint labels, never endpoint IDs, orient it.""" + src_label = graph_view["node_labels"][src] + dst_label = graph_view["node_labels"][dst] + return min(src_label, dst_label), int(edge_label), max(src_label, dst_label) + + @staticmethod + def _node_structure_signatures( + num_nodes: int, + node_labels: Sequence[int], + arcs: Sequence[Tuple[int, int, int]], + ) -> Dict[int, Any]: + """Refine node signatures using labels and labeled undirected neighborhoods. + + The signature is intentionally independent of the input node numbering. + Nodes that remain tied after refinement are structurally equivalent for + the supported labeled simple-graph contract. + """ + neighbors: Dict[int, List[Tuple[int, int]]] = defaultdict(list) + for src, dst, edge_label in arcs: + neighbors[src].append((int(edge_label), dst)) + signatures: Dict[int, int] = { + node: int(node_labels[node]) for node in range(num_nodes) + } + for _ in range(num_nodes): + descriptions = { + node: (int(node_labels[node]), tuple(sorted( + (edge_label, signatures[neighbor]) + for edge_label, neighbor in neighbors[node]))) + for node in range(num_nodes) + } + color_by_description = { + description: color + for color, description in enumerate(sorted(set(descriptions.values()))) + } + refined = { + node: color_by_description[description] + for node, description in descriptions.items() + } + if refined == signatures: + break + signatures = refined + return signatures + + @staticmethod + def _select_structural_node( + graph_view: Dict[str, Any], candidates: Iterable[int]) -> int: + """Select by structure, retaining an arbitrary member only when tied.""" + selected = None + selected_signature = None + for node in candidates: + signature = graph_view["node_signatures"][node] + if selected is None or signature < selected_signature: + selected = node + selected_signature = signature + if selected is None: + raise ValueError("Feuler component has no node eligible for traversal.") + return selected + + def _traversal_signature( + self, + graph_view: Dict[str, Any], + edges: Sequence[Tuple[int, int, int]], + component: Sequence[int], + ) -> Tuple[Any, ...]: + """Compare components without reversible raw node-reference tokens.""" + if not edges: + node = self._select_structural_node(graph_view, component) + return (("node", graph_view["node_labels"][node]),) + signature: List[Any] = [("node", graph_view["node_labels"][edges[0][0]])] + signature.extend( + ("edge", edge_label, graph_view["node_labels"][dst]) + for _, dst, edge_label in edges) + return tuple(signature) + + def _edges_to_tokens( + self, + graph_view: Dict[str, Any], + edges: Sequence[Tuple[int, int, int]], + component: Sequence[int], + ) -> List[int]: + if not edges: + return self._node_tokens( + graph_view, self._select_structural_node(graph_view, component)) + + tokens = self._node_tokens(graph_view, edges[0][0]) + for _, dst, edge_label in edges: + if self.include_edge_tokens: + tokens.append(self.edge_token(edge_label)) + tokens.extend(self._node_tokens(graph_view, dst)) + return tokens + + def _node_tokens(self, graph_view: Dict[str, Any], node: int) -> List[int]: + return [ + self.node_token(graph_view["node_labels"][node]), + self.node_reference_token(node), + ] + + @staticmethod + def _to_list(value: Any): + if hasattr(value, "tolist"): + return value.tolist() + return value + + @staticmethod + def _is_sequence(value: Any) -> bool: + return isinstance(value, Iterable) and not isinstance(value, (str, bytes)) + + def _scalar(self, value: Any) -> int: + value = self._to_list(value) + if self._is_sequence(value): + if len(value) != 1: + raise ValueError( + "Multi-dimensional node and edge features require explicit " + "dataset-specific token encoding.") + value = value[0] + if self._is_sequence(value): + raise ValueError( + "Multi-dimensional node and edge features require explicit " + "dataset-specific token encoding.") + return int(value) + + +class EulerianSerializer(FrequencyGuidedEulerianSerializer): + """Deterministic Eulerian serializer without corpus frequency guidance.""" + + def __init__(self, include_edge_tokens: bool = True, component_sep_token_id: int = -1): + super().__init__( + name="eulerian", + include_edge_tokens=include_edge_tokens, + component_sep_token_id=component_sep_token_id, + ) + + def fit(self, graphs: Sequence[Any]): + self.frequency_map = {} + return self + + def _serialize_component( + self, + graph_view: Dict[str, Any], + component: Sequence[int], + ) -> List[Tuple[int, int, int]]: + component_nodes = set(component) + arcs = [ + (src, dst, label) + for src, dst, label in graph_view["arcs"] + if src in component_nodes and dst in component_nodes + ] + if not arcs: + return [] + + adjacency: Dict[int, List[Tuple[int, int]]] = defaultdict(list) + for src, dst, label in arcs: + adjacency[src].append((dst, label)) + for src in adjacency: + adjacency[src].sort( + key=lambda item: (item[1], graph_view["node_signatures"][item[0]]), + reverse=True, + ) + + start = self._select_structural_node( + graph_view, (node for node in component if adjacency.get(node))) + node_stack = [start] + edge_stack: List[Tuple[int, int, int]] = [] + circuit: List[Tuple[int, int, int]] = [] + while node_stack: + current = node_stack[-1] + if adjacency[current]: + dst, label = adjacency[current].pop() + node_stack.append(dst) + edge_stack.append((current, dst, label)) + else: + node_stack.pop() + if edge_stack: + circuit.append(edge_stack.pop()) + circuit.reverse() + return circuit + + +GraphSerializer = FrequencyGuidedEulerianSerializer diff --git a/gammagl/transforms/graph_tokenizer.py b/gammagl/transforms/graph_tokenizer.py new file mode 100644 index 000000000..43d7fd5dc --- /dev/null +++ b/gammagl/transforms/graph_tokenizer.py @@ -0,0 +1,379 @@ +from dataclasses import dataclass +import hashlib +import random +from typing import Any, Dict, List + +from .base_transform import BaseTransform +from .graph_bpe import GraphBPE +from .graph_serializer import FrequencyGuidedEulerianSerializer + + +@dataclass(frozen=True) +class GraphTokenizerSpecialTokens: + pad_token_id: int = 0 + unk_token_id: int = 1 + mask_token_id: int = 2 + cls_token_id: int = 3 + sep_token_id: int = 4 + node_start_token_id: int = 5 + node_end_token_id: int = 6 + component_sep_token_id: int = 7 + + +@dataclass +class GraphTokenizationResult: + """Tokenized graph sequence and lightweight provenance.""" + + input_ids: List[int] + attention_mask: List[int] + serialized_token_ids: List[int] + metadata: Dict[str, Any] + + +@dataclass +class GraphTokenizerMLMBatch: + """Padded masked-token batch for MLM pretraining.""" + + input_ids: List[List[int]] + attention_mask: List[List[int]] + labels: List[List[int]] + metadata: Dict[str, Any] + + +class GraphTokenizer(BaseTransform): + """Composes graph serialization and BPE into token ids. + + With the default Feuler serializer, inputs are limited to simple + undirected graphs without self-loops or parallel edges. A custom + serializer owns any broader graph-domain contract. + """ + + def __init__( + self, + serializer: FrequencyGuidedEulerianSerializer | None = None, + bpe: GraphBPE | None = None, + special_tokens: GraphTokenizerSpecialTokens | None = None, + add_special_tokens: bool = True, + ): + self.serializer = serializer or FrequencyGuidedEulerianSerializer() + self.bpe = bpe or GraphBPE() + self.special_tokens = special_tokens or GraphTokenizerSpecialTokens() + if not hasattr(self.bpe, "protected_token_ids"): + self.bpe.protected_token_ids = set() + self.bpe.protected_token_ids.add( + self.special_tokens.component_sep_token_id) + self.add_special_tokens = bool(add_special_tokens) + self._fitted = False + self.vocabulary = {} + self.id_to_vocabulary_token = {} + self.fit_graph_ids_hash = None + self.schema_version = 1 + + def fit(self, graphs, graph_ids=None): + graphs = list(graphs) + if not graphs: + raise ValueError("Cannot fit GraphTokenizer on an empty training corpus.") + self.fit_graph_ids_hash = None + if graph_ids is not None: + graph_ids = list(graph_ids) + if len(graph_ids) != len(graphs): + raise ValueError("graph_ids must align with the training graphs.") + self.fit_graph_ids_hash = hashlib.sha256( + repr(tuple(graph_ids)).encode("utf-8")).hexdigest() + self.serializer.fit(graphs) + serialized_sequences = [ + self._normalize_serialized_tokens(self.serializer.serialize(graph).token_ids) + for graph in graphs + ] + self.bpe.fit(serialized_sequences) + # Match the official protocol: BPE IDs are an intermediate alphabet, + # then a frozen train-built contiguous vocabulary maps them to model IDs. + # This permits inductive val/test encoding through UNK without changing + # BPE rules or admitting evaluation graphs into vocabulary fitting. + base_ids = {int(token) for sequence in serialized_sequences for token in sequence} + merge_ids = {int(rule[2]) for rule in self.bpe.codebook.merge_rules} + special_ids = set(vars(self.special_tokens).values()) + token_ids = sorted((base_ids | merge_ids) - special_ids) + start = self._reserved_token_count() + self.vocabulary = {token: start + index for index, token in enumerate(token_ids)} + self.id_to_vocabulary_token = { + vocab_id: token for token, vocab_id in self.vocabulary.items()} + self._fitted = True + return self + + def encode_tokens(self, token_ids: List[int]) -> List[int]: + encoded = self._vocabulary_encode(self.bpe.encode(token_ids)) + if not self.add_special_tokens: + return encoded + return [self.special_tokens.cls_token_id, *encoded, self.special_tokens.sep_token_id] + + @property + def max_token_id(self) -> int: + return max( + max(self.vocabulary.values(), default=self._reserved_token_count() - 1), + self.special_tokens.pad_token_id, + self.special_tokens.unk_token_id, + self.special_tokens.mask_token_id, + self.special_tokens.cls_token_id, + self.special_tokens.sep_token_id, + self.special_tokens.node_start_token_id, + self.special_tokens.node_end_token_id, + self.special_tokens.component_sep_token_id, + ) + + def validate_model_vocab(self, vocab_size: int) -> None: + required_vocab_size = self.max_token_id + 1 + if int(vocab_size) < required_vocab_size: + raise ValueError( + f"Model vocab_size={vocab_size} is smaller than the tokenizer " + f"requirement {required_vocab_size}.") + + def encode_graph(self, graph: Any) -> GraphTokenizationResult: + self._require_fitted() + serialized = self.serializer.serialize(graph) + serialized_token_ids = self._normalize_serialized_tokens(serialized.token_ids) + input_ids = self.encode_tokens(serialized_token_ids) + return GraphTokenizationResult( + input_ids=input_ids, + attention_mask=[1] * len(input_ids), + serialized_token_ids=serialized_token_ids, + metadata={ + "serializer": serialized.metadata, + "bpe": { + "num_merge_rules": len(self.bpe.codebook.merge_rules), + "vocab_size": self.bpe.codebook.vocab_size, + }, + "required_vocab_size": self.max_token_id + 1, + "add_special_tokens": self.add_special_tokens, + "tokenizer_schema_version": self.schema_version, + "fit_graph_ids_hash": self.fit_graph_ids_hash, + }, + ) + + def decode_graph(self, encoding) -> Dict[str, Any]: + """Invert BPE and reconstruct a graph from an encoded token sequence.""" + self._require_fitted() + if isinstance(encoding, GraphTokenizationResult): + input_ids = encoding.input_ids + else: + input_ids = list(encoding) + if self.add_special_tokens: + if ( + len(input_ids) < 2 + or input_ids[0] != self.special_tokens.cls_token_id + or input_ids[-1] != self.special_tokens.sep_token_id): + raise ValueError("Encoded graph must begin with CLS and end with SEP tokens.") + input_ids = input_ids[1:-1] + serialized = self.bpe.decode([ + self.id_to_vocabulary_token.get(token, token) + for token in input_ids + ]) + return self.serializer.deserialize( + self._denormalize_serialized_tokens(serialized)) + + def batch_encode_graphs(self, graphs) -> List[GraphTokenizationResult]: + self._require_fitted() + graphs = list(graphs) + serialized_results = [ + self.serializer.serialize(graph) for graph in graphs + ] + serialized_token_ids = [ + self._normalize_serialized_tokens(result.token_ids) + for result in serialized_results + ] + encoded_sequences = [ + self._vocabulary_encode(sequence) + for sequence in self.bpe.batch_encode(serialized_token_ids) + ] + results = [] + for serialized, normalized, encoded in zip( + serialized_results, serialized_token_ids, encoded_sequences): + encoded = list(map(int, encoded)) + if self.add_special_tokens: + encoded = [ + self.special_tokens.cls_token_id, + *encoded, + self.special_tokens.sep_token_id, + ] + results.append(GraphTokenizationResult( + input_ids=encoded, + attention_mask=[1] * len(encoded), + serialized_token_ids=normalized, + metadata={ + "serializer": serialized.metadata, + "bpe": { + "num_merge_rules": len(self.bpe.codebook.merge_rules), + "vocab_size": self.bpe.codebook.vocab_size, + }, + "required_vocab_size": self.max_token_id + 1, + "add_special_tokens": self.add_special_tokens, + "tokenizer_schema_version": self.schema_version, + "fit_graph_ids_hash": self.fit_graph_ids_hash, + }, + )) + return results + + def _normalize_serialized_tokens(self, token_ids: List[int]) -> List[int]: + reserved_token_count = self._reserved_token_count() + component_sep_token_id = getattr( + self.serializer, + "component_sep_token_id", + self.special_tokens.component_sep_token_id, + ) + normalized = [] + for token in token_ids: + token = int(token) + if token == component_sep_token_id: + normalized.append(self.special_tokens.component_sep_token_id) + else: + normalized.append(token + reserved_token_count) + return normalized + + def _denormalize_serialized_tokens(self, token_ids: List[int]) -> List[int]: + reserved_token_count = self._reserved_token_count() + denormalized = [] + for token in token_ids: + token = int(token) + if token == self.special_tokens.component_sep_token_id: + denormalized.append(self.serializer.component_sep_token_id) + elif token < reserved_token_count: + raise ValueError(f"Encoded graph contains reserved token id {token}.") + else: + denormalized.append(token - reserved_token_count) + return denormalized + + def _reserved_token_count(self) -> int: + return max( + self.special_tokens.pad_token_id, + self.special_tokens.unk_token_id, + self.special_tokens.mask_token_id, + self.special_tokens.cls_token_id, + self.special_tokens.sep_token_id, + self.special_tokens.node_start_token_id, + self.special_tokens.node_end_token_id, + self.special_tokens.component_sep_token_id, + ) + 1 + + def _vocabulary_encode(self, token_ids: List[int]) -> List[int]: + special_ids = set(vars(self.special_tokens).values()) + return [ + int(token) if int(token) in special_ids + else self.vocabulary.get(int(token), self.special_tokens.unk_token_id) + for token in token_ids + ] + + def build_mlm_batch( + self, + graphs, + max_length: int, + mask_prob: float = 0.15, + seed: int = 0, + ) -> GraphTokenizerMLMBatch: + return self.build_mlm_batch_from_token_sequences( + [self.encode_graph(graph).input_ids for graph in graphs], + max_length=max_length, + mask_prob=mask_prob, + seed=seed, + ) + + def pad_token_sequences(self, token_sequences, max_length: int): + if max_length <= 0: + raise ValueError("max_length must be positive.") + if any(len(sequence) > max_length for sequence in token_sequences): + raise ValueError( + "Token sequence exceeds max_length; truncation would discard graph structure.") + input_ids = [ + [int(token) for token in sequence] + for sequence in token_sequences + ] + attention_mask = [ + [1] * len(sequence) + [0] * (max_length - len(sequence)) + for sequence in input_ids + ] + padded_ids = [ + sequence + [self.special_tokens.pad_token_id] * (max_length - len(sequence)) + for sequence in input_ids + ] + return padded_ids, attention_mask + + def build_mlm_batch_from_token_sequences( + self, + token_sequences, + max_length: int, + mask_prob: float = 0.15, + seed: int = 0, + ) -> GraphTokenizerMLMBatch: + padded_ids, attention_mask = self.pad_token_sequences( + token_sequences, max_length=max_length) + masked_ids, labels = self.mask_token_sequences( + padded_ids, mask_prob=mask_prob, seed=seed) + num_mlm_labels = sum(1 for row in labels for value in row if value != -100) + return GraphTokenizerMLMBatch( + input_ids=masked_ids, + attention_mask=attention_mask, + labels=labels, + metadata={ + "mask_prob": float(mask_prob), + "max_length": int(max_length), + "num_graphs": len(padded_ids), + "num_mlm_labels": num_mlm_labels, + }, + ) + + def mask_token_sequences( + self, input_ids: List[List[int]], mask_prob: float, seed: int): + rng = random.Random(seed) + special_ids = { + self.special_tokens.pad_token_id, + self.special_tokens.cls_token_id, + self.special_tokens.sep_token_id, + self.special_tokens.component_sep_token_id, + } + masked_ids = [] + labels = [] + for sequence in input_ids: + masked_sequence = [] + label_sequence = [] + for token in sequence: + token = int(token) + if token in special_ids or rng.random() >= mask_prob: + masked_sequence.append(token) + label_sequence.append(-100) + else: + masked_sequence.append(self.special_tokens.mask_token_id) + label_sequence.append(token) + masked_ids.append(masked_sequence) + labels.append(label_sequence) + if ( + float(mask_prob) > 0 + and not any(label != -100 for row in labels for label in row)): + for row_index, sequence in enumerate(input_ids): + for column_index, token in enumerate(sequence): + token = int(token) + if token not in special_ids: + masked_ids[row_index][column_index] = ( + self.special_tokens.mask_token_id) + labels[row_index][column_index] = token + return masked_ids, labels + return masked_ids, labels + + def _mask_input_ids(self, input_ids: List[List[int]], mask_prob: float, seed: int): + return self.mask_token_sequences(input_ids, mask_prob=mask_prob, seed=seed) + + def _require_fitted(self) -> None: + if not self._fitted: + raise RuntimeError("GraphTokenizer must be fit on training graphs before encoding.") + + def __call__(self, data: Any): + encoding = self.encode_graph(data) + if isinstance(data, dict): + data["input_ids"] = encoding.input_ids + data["attention_mask"] = encoding.attention_mask + data["serialized_token_ids"] = encoding.serialized_token_ids + data["graph_tokenizer_metadata"] = encoding.metadata + else: + setattr(data, "input_ids", encoding.input_ids) + setattr(data, "attention_mask", encoding.attention_mask) + setattr(data, "serialized_token_ids", encoding.serialized_token_ids) + setattr(data, "graph_tokenizer_metadata", encoding.metadata) + return data diff --git a/setup.py b/setup.py index 251e1debc..5e311f2d3 100644 --- a/setup.py +++ b/setup.py @@ -149,6 +149,18 @@ def load_ops_extensions(): # Start to include cuda ops, if no cuda found, will only compile cpu ops def load_extensions(): extensions = load_mpops_extensions() + load_ops_extensions() + try: + import pybind11 + except ImportError: + # The rest of GammaGL remains installable without the optional BPE + # accelerator; GraphBPE auto mode uses its canonical Python path. + return extensions + extensions.append(PyCppExtension( + name='third_party.graph_bpe_cpp._graph_bpe', + sources=[osp.join('third_party', 'graph_bpe_cpp', '_graph_bpe.cpp')], + include_dirs=[pybind11.get_include()], + extra_compile_args=['-std=c++17'], + )) return extensions @@ -194,6 +206,16 @@ def load_extensions(): 'nbsphinx', ], 'defog': ['rdkit', 'networkx'], + 'graph-tokenizer-paper': [ + # Paper protocol runtime; not a GammaGL core dependency. + 'torch>=2.1', + 'dgl>=2.4', + 'torch-geometric>=2.4', + 'ogb>=1.3.6', + # Pinned GTE checkpoint/config acquisition and safetensors conversion. + 'huggingface-hub>=0.20', + 'safetensors>=0.4', + ], 'llm': [ 'torch>=2.1', 'transformers>=4.31', diff --git a/tests/data/test_dataset.py b/tests/data/test_dataset.py index 1c6552f57..a0700295e 100644 --- a/tests/data/test_dataset.py +++ b/tests/data/test_dataset.py @@ -7,6 +7,7 @@ import pytest from gammagl.datasets.ppi import PPI +from gammagl.data.dataset import Dataset # dataset record to avoid downloading repeatedly @@ -22,3 +23,20 @@ def test_dataset(): assert len(dataset1) == 20 assert len(dataset2) == 20 + + +def test_torch_dataset_loader_restores_graph_objects(monkeypatch): + torch = __import__('torch') + observed = {} + + def fake_load(path, **kwargs): + observed['path'] = path + observed.update(kwargs) + return ('graph', None) + + monkeypatch.setattr(torch, 'load', fake_load) + dataset = object.__new__(Dataset) + + assert dataset.load_data('trusted_graphs.pt') == ('graph', None) + assert observed['path'] == 'trusted_graphs.pt' + assert observed['weights_only'] is False diff --git a/tests/datasets/test_graph_tokenizer_dataset_download.py b/tests/datasets/test_graph_tokenizer_dataset_download.py new file mode 100644 index 000000000..6be03a2c1 --- /dev/null +++ b/tests/datasets/test_graph_tokenizer_dataset_download.py @@ -0,0 +1,214 @@ +import importlib.util +import hashlib +import json +import pickle +import threading +import time +import zipfile +from pathlib import Path + +import pytest + + +def _load_download_module(): + module_path = ( + Path(__file__).resolve().parents[2] + / "gammagl" + / "datasets" + / "_graph_tokenizer_download.py" + ) + spec = importlib.util.spec_from_file_location( + "graph_tokenizer_download_under_test", module_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +@pytest.mark.parametrize( + ("dataset_name", "data_filename"), + ( + ("qm9", "data.pkl"), + ("ogbg-molhiv", "data.pkl"), + ("peptides-struct", "data.pkl.gz"), + ), +) +def test_materialize_dataset_from_cached_official_bundle( + tmp_path, monkeypatch, dataset_name, data_filename): + module = _load_download_module() + bundle_root = tmp_path / "release" + source_dir = bundle_root / "data" / dataset_name + source_dir.mkdir(parents=True) + (source_dir / data_filename).write_bytes(pickle.dumps([dataset_name])) + for split, indices in {"train": [0], "val": [], "test": []}.items(): + (source_dir / f"{split}_index.json").write_text( + json.dumps(indices), encoding="utf-8") + + monkeypatch.setenv(module.DATA_BUNDLE_ENV, str(bundle_root)) + raw_dir = tmp_path / "datasets" / dataset_name / "raw" + + copied = module.materialize_paper_dataset( + dataset_name=dataset_name, + aliases=(dataset_name.replace("-", "_"),), + raw_dir=raw_dir, + cache_root=tmp_path / "datasets", + allow_download=True, + ) + + assert copied == raw_dir + assert (raw_dir / data_filename).read_bytes() == (source_dir / data_filename).read_bytes() + assert json.loads((raw_dir / "train_index.json").read_text(encoding="utf-8")) == [0] + + +def test_materialize_dataset_rejects_bundle_without_official_splits(tmp_path, monkeypatch): + module = _load_download_module() + bundle_root = tmp_path / "release" + source_dir = bundle_root / "data" / "qm9" + source_dir.mkdir(parents=True) + (source_dir / "data.pkl").write_bytes(pickle.dumps([])) + monkeypatch.setenv(module.DATA_BUNDLE_ENV, str(bundle_root)) + + with pytest.raises(FileNotFoundError, match="official split"): + module.materialize_paper_dataset( + dataset_name="qm9", + aliases=(), + raw_dir=tmp_path / "qm9" / "raw", + cache_root=tmp_path, + allow_download=True, + ) + + +def test_download_bundle_uses_actual_file_sha256(tmp_path, monkeypatch): + module = _load_download_module() + archive = tmp_path / "official.tar.gz" + archive.write_bytes(b"official release bytes") + expected = hashlib.sha256(b"official release bytes").hexdigest() + monkeypatch.setattr(module, "PAPER_DATA_BUNDLE_SHA256", expected) + + assert module.sha256_file(archive) == expected + assert module._verify_official_bundle(archive, remove_on_failure=True) == expected + + +def test_cached_bundle_checksum_mismatch_is_deleted_and_hard_fails(tmp_path, monkeypatch): + module = _load_download_module() + archive = (tmp_path / ".graph_tokenizer_release" + / module.PAPER_DATA_BUNDLE_FILENAME) + archive.parent.mkdir(parents=True) + archive.write_bytes(b"tampered release bytes") + monkeypatch.setattr(module, "PAPER_DATA_BUNDLE_SHA256", "0" * 64) + + with pytest.raises(RuntimeError, match="SHA-256 mismatch"): + module._resolve_bundle_source(tmp_path) + assert archive.exists() + + +def test_training_materialization_requires_prepared_bundle(tmp_path, monkeypatch): + module = _load_download_module() + monkeypatch.delenv(module.DATA_BUNDLE_ENV, raising=False) + + with pytest.raises(FileNotFoundError, match="prepared by the single data-preparation"): + module.materialize_paper_dataset( + dataset_name="qm9", + aliases=(), + raw_dir=tmp_path / "qm9" / "raw", + cache_root=tmp_path, + ) + + +def test_preparation_installs_verified_bundle_atomically(tmp_path, monkeypatch): + module = _load_download_module() + payload = b"official release bytes" + monkeypatch.setattr(module, "PAPER_DATA_BUNDLE_SHA256", hashlib.sha256(payload).hexdigest()) + + def fake_download(_file_id, destination, filename): + path = Path(destination) / filename + path.write_bytes(payload) + return str(path) + + import gammagl.data.download as download + monkeypatch.setattr(download, "download_google_url", fake_download) + + archive = module._resolve_bundle_source(tmp_path, allow_download=True) + + assert archive.name == module.PAPER_DATA_BUNDLE_FILENAME + assert archive.read_bytes() == payload + assert not list(archive.parent.glob(f".{archive.name}.*.tmp")) + + +def test_concurrent_preparation_has_one_final_bundle_writer(tmp_path, monkeypatch): + module = _load_download_module() + payload = b"official release bytes" + monkeypatch.setattr(module, "PAPER_DATA_BUNDLE_SHA256", hashlib.sha256(payload).hexdigest()) + calls = [] + calls_lock = threading.Lock() + + def fake_download(_file_id, destination, filename): + with calls_lock: + calls.append(filename) + time.sleep(0.05) + path = Path(destination) / filename + path.write_bytes(payload) + return str(path) + + import gammagl.data.download as download + monkeypatch.setattr(download, "download_google_url", fake_download) + results = [] + failures = [] + + def prepare(): + try: + results.append(module._resolve_bundle_source(tmp_path, allow_download=True)) + except Exception as error: # pragma: no cover - asserted below + failures.append(error) + + workers = [threading.Thread(target=prepare) for _ in range(2)] + for worker in workers: + worker.start() + for worker in workers: + worker.join() + + assert not failures + assert len(calls) == 1 + assert len(results) == 2 + assert results[0] == results[1] + assert results[0].read_bytes() == payload + + +def test_failed_downloader_keeps_existing_verified_shared_bundle(tmp_path, monkeypatch): + module = _load_download_module() + payload = b"official release bytes" + archive = tmp_path / ".graph_tokenizer_release" / module.PAPER_DATA_BUNDLE_FILENAME + archive.parent.mkdir(parents=True) + archive.write_bytes(payload) + monkeypatch.setattr(module, "PAPER_DATA_BUNDLE_SHA256", hashlib.sha256(payload).hexdigest()) + + import gammagl.data.download as download + monkeypatch.setattr( + download, "download_google_url", + lambda *_args, **_kwargs: (_ for _ in ()).throw(RuntimeError("network failed")), + ) + + assert module._resolve_bundle_source(tmp_path, allow_download=True) == archive + assert archive.read_bytes() == payload + + +def test_explicit_archive_checksum_mismatch_hard_fails(tmp_path, monkeypatch): + module = _load_download_module() + archive = tmp_path / "untrusted.tar.gz" + archive.write_bytes(b"not the official release") + monkeypatch.setenv(module.DATA_BUNDLE_ENV, str(archive)) + monkeypatch.setattr(module, "PAPER_DATA_BUNDLE_SHA256", "f" * 64) + + with pytest.raises(RuntimeError, match="SHA-256 mismatch"): + module._resolve_bundle_source(tmp_path, allow_download=True) + assert archive.exists() + + +def test_bundle_extraction_rejects_path_traversal(tmp_path): + module = _load_download_module() + archive = tmp_path / "unsafe.zip" + with zipfile.ZipFile(archive, "w") as handle: + handle.writestr("../outside", b"unsafe") + + with pytest.raises(ValueError, match="Unsafe path"): + module._extract_zip(archive, tmp_path / "extracted") + assert not (tmp_path / "outside").exists() diff --git a/tests/datasets/test_molecule_benchmark_datasets.py b/tests/datasets/test_molecule_benchmark_datasets.py new file mode 100644 index 000000000..6f7c64776 --- /dev/null +++ b/tests/datasets/test_molecule_benchmark_datasets.py @@ -0,0 +1,196 @@ +import gzip +import json +import pickle + +import pytest +import tensorlayerx as tlx + +from gammagl.datasets import OGBGMolHIV, PeptidesStruct, QM9 + + +def _write_preprocessed_raw(root, folder_name, samples, splits, compressed=False): + raw_dir = root / folder_name / "raw" + raw_dir.mkdir(parents=True) + data_name = "data.pkl.gz" if compressed else "data.pkl" + opener = gzip.open if compressed else open + with opener(raw_dir / data_name, "wb") as f: + pickle.dump(samples, f) + for split_name, indices in splits.items(): + (raw_dir / f"{split_name}_index.json").write_text(json.dumps(indices), encoding="utf-8") + + +def test_qm9_loads_preprocessed_raw_files(tmp_path): + properties = { + "mu": 0.1, + "alpha": 0.2, + "homo": 0.3, + "lumo": 0.4, + "gap": 0.5, + "r2": 0.6, + "zpve": 0.7, + "u0": 0.8, + "u298": 0.9, + "h298": 1.0, + "g298": 1.1, + "cv": 1.2, + "u0_atom": 1.3, + "u298_atom": 1.4, + "h298_atom": 1.5, + "g298_atom": 1.6, + } + samples = [ + { + "edge_index": [[0], [1]], + "x": [6, 8], + "edge_attr": [1], + "properties": properties, + }, + { + "edge_index": [[0], [1]], + "x": [7, 6], + "edge_attr": [2], + "properties": {key: value + 1.0 for key, value in properties.items()}, + }, + ] + _write_preprocessed_raw(tmp_path, "qm9", samples, {"train": [0], "val": [1], "test": []}) + + dataset = QM9(root=str(tmp_path)) + + assert len(dataset) == 2 + assert dataset.num_tasks == 16 + assert dataset.metric == "mae" + assert dataset.get_idx_split() == {"train": [0], "val": [1], "test": []} + assert len(dataset.get_split("train")) == 1 + assert tlx.convert_to_numpy(dataset[0].edge_index).tolist() == [[0], [1]] + assert tlx.convert_to_numpy(dataset[0].y).tolist()[0][:3] == pytest.approx( + [0.1, 0.2, 0.3]) + assert tlx.convert_to_numpy(dataset[0].x).tolist() == [13, 17] + assert tlx.convert_to_numpy(dataset[0].edge_attr).tolist() == [2] + + +def test_qm9_converts_official_attr_columns_to_paper_token_ids(tmp_path): + properties = { + name: float(index) for index, name in enumerate(QM9.label_keys) + } + node_attr = [ + [0, 0, 0, 0, 0, 6], + [0, 0, 0, 0, 0, 8], + ] + edge_attr = [[0, 0, 1, 0]] + samples = [{ + "edge_index": [[0], [1]], + "attr": node_attr, + "edge_attr": edge_attr, + "properties": properties, + }] + _write_preprocessed_raw( + tmp_path, "qm9", samples, + {"train": [0], "val": [], "test": []}, + ) + + dataset = QM9(root=str(tmp_path)) + + assert tlx.convert_to_numpy(dataset[0].x).tolist() == [13, 17] + assert tlx.convert_to_numpy(dataset[0].edge_attr).tolist() == [6] + + +def test_qm9_loads_data_prepared_from_released_bundle(tmp_path, monkeypatch): + properties = { + name: float(index) + for index, name in enumerate(QM9.label_keys) + } + bundle_dir = tmp_path / "paper-release" / "data" / "qm9" + bundle_dir.mkdir(parents=True) + with (bundle_dir / "data.pkl").open("wb") as handle: + pickle.dump([{ + "edge_index": [[0], [1]], + "x": [6, 8], + "edge_attr": [1], + "properties": properties, + }], handle) + for split, indices in {"train": [0], "val": [], "test": []}.items(): + (bundle_dir / f"{split}_index.json").write_text( + json.dumps(indices), encoding="utf-8") + monkeypatch.setenv( + "GAMMAGL_GRAPH_TOKENIZER_DATA_BUNDLE", + str(tmp_path / "paper-release"), + ) + + from gammagl.datasets._graph_tokenizer_download import materialize_paper_dataset + materialize_paper_dataset( + dataset_name="qm9", + aliases=("qm9",), + raw_dir=tmp_path / "datasets" / "qm9" / "raw", + cache_root=tmp_path / "datasets", + allow_download=True, + ) + + dataset = QM9(root=str(tmp_path / "datasets")) + + assert len(dataset) == 1 + assert dataset.get_idx_split() == {"train": [0], "val": [], "test": []} + + +def test_qm9_rejects_missing_paper_target_property(tmp_path): + samples = [ + { + "edge_index": [[0], [1]], + "x": [6, 8], + "edge_attr": [1], + "properties": {"mu": 0.1}, + } + ] + _write_preprocessed_raw(tmp_path, "qm9", samples, {"train": [0], "val": [], "test": []}) + + with pytest.raises(ValueError, match="16 QM9 properties"): + QM9(root=str(tmp_path)) + + +def test_dataset_rejects_overlapping_split_indices(tmp_path): + samples = [ + ({"edges": [[0], [1]], "node_type_ids": [6, 8], "edge_type_ids": [1]}, [1]), + ({"edges": [[0], [1]], "node_type_ids": [6, 6], "edge_type_ids": [1]}, [0]), + ] + _write_preprocessed_raw(tmp_path, "ogbg-molhiv", samples, {"train": [0], "val": [0], "test": [1]}) + + with pytest.raises(ValueError, match="overlaps"): + OGBGMolHIV(root=str(tmp_path)) + + +def test_ogbg_molhiv_loads_alias_directory_and_label(tmp_path): + samples = [ + ({"edges": [[0, 1], [1, 0]], "node_type_ids": [6, 6], "edge_type_ids": [1, 1]}, [0]), + ({"edges": [[0], [1]], "node_type_ids": [6, 8], "edge_type_ids": [2]}, [1]), + ] + _write_preprocessed_raw(tmp_path, "ogbg-molhiv", samples, {"train": [0], "val": [], "test": [1]}) + + dataset = OGBGMolHIV(root=str(tmp_path)) + + assert len(dataset) == 2 + assert dataset.name == "ogbg-molhiv" + assert dataset.num_tasks == 1 + assert dataset.metric == "rocauc" + assert dataset.get_idx_split()["test"] == [1] + assert tlx.convert_to_numpy(dataset[1].y).tolist() == [[1.0]] + + +def test_peptides_struct_loads_compressed_raw_files(tmp_path): + samples = [ + ( + { + "edges": [[0, 1], [1, 0]], + "node_token_ids": [[5], [9]], + "edge_token_ids": [[2], [2]], + }, + [float(i) for i in range(11)], + ) + ] + _write_preprocessed_raw(tmp_path, "peptides-struct", samples, {"train": [0], "val": [], "test": []}, compressed=True) + + dataset = PeptidesStruct(root=str(tmp_path)) + + assert len(dataset) == 1 + assert dataset.num_tasks == 11 + assert dataset.metric == "average_mae" + assert tlx.convert_to_numpy(dataset[0].x).tolist() == [5, 9] + assert tlx.convert_to_numpy(dataset[0].y).shape == (1, 11) diff --git a/tests/models/test_graph_gte_pretrained.py b/tests/models/test_graph_gte_pretrained.py new file mode 100644 index 000000000..c610e6bfa --- /dev/null +++ b/tests/models/test_graph_gte_pretrained.py @@ -0,0 +1,193 @@ +import hashlib +import importlib.util +import os +import sys +from pathlib import Path + +import numpy as np +import pytest +import tensorlayerx as tlx + + +torch = pytest.importorskip("torch") +safetensors_torch = pytest.importorskip("safetensors.torch") + + +ROOT = Path(__file__).resolve().parents[2] + + +def _load(name, relative_path): + spec = importlib.util.spec_from_file_location(name, ROOT / relative_path) + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +def _modules(): + _load("graph_bert", "gammagl/models/graph_bert.py") + pretrained = _load("graph_gte_pretrained", "gammagl/models/graph_gte_pretrained.py") + graph_gte = _load("graph_gte_pretrained_test_model", "gammagl/models/graph_gte.py") + return graph_gte.GraphGTE, pretrained + + +def _tiny_checkpoint(model): + layer = model.encoder_layers[0] + values = { + "new.embeddings.token_type_embeddings.weight": model.token_type_embeddings.embeddings, + "new.embeddings.LayerNorm.weight": model.embedding_norm.gamma, + "new.embeddings.LayerNorm.bias": model.embedding_norm.beta, + "new.encoder.layer.0.attention.qkv_proj.weight": layer.attention.qkv_proj.weights, + "new.encoder.layer.0.attention.qkv_proj.bias": layer.attention.qkv_proj.biases, + "new.encoder.layer.0.attention.o_proj.weight": layer.attention.out_proj.weights, + "new.encoder.layer.0.attention.o_proj.bias": layer.attention.out_proj.biases, + "new.encoder.layer.0.attn_ln.weight": layer.attention_norm.gamma, + "new.encoder.layer.0.attn_ln.bias": layer.attention_norm.beta, + "new.encoder.layer.0.mlp.up_gate_proj.weight": layer.mlp.up_gate_proj.weights, + "new.encoder.layer.0.mlp.down_proj.weight": layer.mlp.down_proj.weights, + "new.encoder.layer.0.mlp.down_proj.bias": layer.mlp.down_proj.biases, + "new.encoder.layer.0.mlp_ln.weight": layer.mlp_norm.gamma, + "new.encoder.layer.0.mlp_ln.bias": layer.mlp_norm.beta, + "new.embeddings.word_embeddings.weight": torch.zeros((7, 16)), + "classifier.weight": torch.zeros((2, 16)), + "classifier.bias": torch.zeros((2,)), + } + return {key: torch.arange(value.numel(), dtype=torch.float32).reshape(value.shape) + if key.startswith("new.encoder") else torch.ones_like(value.detach()) + for key, value in values.items()} + + +def test_sha256_file_is_content_hash_and_detects_mismatch(tmp_path): + _, pretrained = _modules() + checkpoint = tmp_path / "model.safetensors" + checkpoint.write_bytes(b"strict-checksum") + assert pretrained.sha256_file(checkpoint) == hashlib.sha256(b"strict-checksum").hexdigest() + assert pretrained.sha256_file(checkpoint) != pretrained.GTE_CHECKPOINT_SHA256 + + +def test_download_rejects_checkpoint_sha256_mismatch(tmp_path, monkeypatch): + _, pretrained = _modules() + checkpoint = tmp_path / "model.safetensors" + checkpoint.write_bytes(b"tampered") + import huggingface_hub + monkeypatch.setattr(huggingface_hub, "hf_hub_download", lambda *args, **kwargs: str(checkpoint)) + with pytest.raises(RuntimeError, match="SHA-256 mismatch"): + pretrained.download_official_gte_checkpoint() + + +def test_explicit_converter_has_full_coverage_and_parameter_equality(tmp_path): + GraphGTE, pretrained = _modules() + source_model = GraphGTE(vocab_size=7, output_dim=2, hidden_size=16, + num_hidden_layers=1, num_attention_heads=4, + intermediate_size=32, max_position_embeddings=16) + tensors = _tiny_checkpoint(source_model) + checkpoint = tmp_path / "tiny.safetensors" + safetensors_torch.save_file(tensors, str(checkpoint)) + model = GraphGTE(vocab_size=7, output_dim=2, hidden_size=16, + num_hidden_layers=1, num_attention_heads=4, + intermediate_size=32, max_position_embeddings=16) + report = pretrained.load_pretrained_encoder(model, checkpoint) + assert report["coverage"] == 1.0 + assert report["missing"] == [] + assert report["shape_mismatches"] == [] + assert report["unexpected"] == [] + assert "text-vocabulary embedding" in report["ignored"]["new.embeddings.word_embeddings.weight"] + assert not np.array_equal( + tlx.convert_to_numpy(model.token_embeddings.embeddings), + tensors["new.embeddings.word_embeddings.weight"].numpy()) + layer = model.encoder_layers[0] + checks = { + "new.encoder.layer.0.attention.qkv_proj.weight": layer.attention.qkv_proj.weights, + "new.encoder.layer.0.attention.o_proj.weight": layer.attention.out_proj.weights, + "new.encoder.layer.0.mlp.up_gate_proj.weight": layer.mlp.up_gate_proj.weights, + "new.encoder.layer.0.mlp_ln.weight": layer.mlp_norm.gamma, + } + for key, target in checks.items(): + np.testing.assert_array_equal(tlx.convert_to_numpy(target), tensors[key].numpy()) + + +def test_converter_rejects_shape_mismatch(tmp_path): + GraphGTE, pretrained = _modules() + source = GraphGTE(vocab_size=7, output_dim=2, hidden_size=16, + num_hidden_layers=1, num_attention_heads=4, + intermediate_size=32, max_position_embeddings=16) + tensors = _tiny_checkpoint(source) + tensors["new.encoder.layer.0.attention.qkv_proj.weight"] = torch.zeros((1, 1)) + checkpoint = tmp_path / "bad.safetensors" + safetensors_torch.save_file(tensors, str(checkpoint)) + with pytest.raises(RuntimeError, match="Incomplete GTE encoder conversion"): + pretrained.load_pretrained_encoder(source, checkpoint) + + +def test_official_checkpoint_converter_and_hf_encoder_equivalence(): + """Real pinned GTE checkpoint check; opt in with GTE_CHECKPOINT_PATH.""" + checkpoint_path = os.environ.get("GTE_CHECKPOINT_PATH") + if not checkpoint_path: + pytest.skip("set GTE_CHECKPOINT_PATH to run the pinned-GTE integration test") + transformers = pytest.importorskip("transformers") + GraphGTE, pretrained = _modules() + assert pretrained.sha256_file(checkpoint_path) == pretrained.GTE_CHECKPOINT_SHA256 + model = GraphGTE(vocab_size=19, output_dim=2) + report = pretrained.load_pretrained_encoder(model, checkpoint_path) + assert report["coverage"] == 1.0 + assert report["missing"] == [] + assert report["shape_mismatches"] == [] + + from safetensors import safe_open + with safe_open(checkpoint_path, framework="pt", device="cpu") as source: + equality_checks = { + "new.encoder.layer.0.attention.qkv_proj.weight": + model.encoder_layers[0].attention.qkv_proj.weights, + "new.encoder.layer.0.attention.o_proj.weight": + model.encoder_layers[0].attention.out_proj.weights, + "new.encoder.layer.6.mlp.up_gate_proj.weight": + model.encoder_layers[6].mlp.up_gate_proj.weights, + "new.encoder.layer.11.mlp_ln.weight": model.encoder_layers[11].mlp_norm.gamma, + } + for key, target in equality_checks.items(): + np.testing.assert_array_equal( + source.get_tensor(key).numpy(), tlx.convert_to_numpy(target)) + + reference = transformers.AutoModel.from_pretrained( + pretrained.GTE_MODEL_ID, revision=pretrained.GTE_REVISION, + trust_remote_code=True, local_files_only=True).eval() + model.set_eval() + torch.manual_seed(11) + inputs = torch.randn(1, 5, 768) + mask = torch.tensor([[1, 1, 1, 1, 0]], dtype=torch.long) + with torch.no_grad(): + reference_outputs = reference( + inputs_embeds=inputs, attention_mask=mask, + output_hidden_states=True, return_dict=True).hidden_states + _, tlx_outputs = model.encode_embeddings( + inputs, mask, return_hidden_states=True) + errors = [float(np.max(np.abs( + reference_output.detach().numpy() - tlx.convert_to_numpy(tlx_output)))) + for reference_output, tlx_output in zip(reference_outputs, tlx_outputs)] + assert len(errors) == 13 + assert errors[1] < 5e-5 # layer 0 + assert errors[7] < 5e-5 # middle layer (hidden_states includes embeddings) + assert errors[-1] < 5e-5 # final layer + + +def test_official_from_pretrained_smoke_updates_encoder(): + if not os.environ.get("GTE_CHECKPOINT_PATH"): + pytest.skip("set GTE_CHECKPOINT_PATH to run the pinned-GTE integration test") + GraphGTE, _ = _modules() + model = GraphGTE.from_pretrained(vocab_size=19, output_dim=2) + manifest = model.manifest() + assert type(model).__name__ == "GraphGTE" + assert manifest["pretrained"] is True + assert manifest["encoder_coverage"] == 1.0 + parameter = model.encoder_layers[0].attention.qkv_proj.weights + before = parameter.detach().clone() + optimizer = torch.optim.SGD(model.trainable_weights, lr=1e-4) + output = model( + torch.tensor([[2, 3, 4, 1]], dtype=torch.long), + torch.tensor([[1, 1, 1, 0]], dtype=torch.long), task="supervised") + loss = (output * output).mean() + optimizer.zero_grad() + loss.backward() + optimizer.step() + assert parameter.grad is not None + assert not torch.equal(before, parameter.detach()) diff --git a/tests/models/test_graph_tokenizer_paper_protocol.py b/tests/models/test_graph_tokenizer_paper_protocol.py new file mode 100644 index 000000000..6b5b4db4f --- /dev/null +++ b/tests/models/test_graph_tokenizer_paper_protocol.py @@ -0,0 +1,1436 @@ +import importlib.util +import sys +from pathlib import Path +from types import SimpleNamespace + +import pytest + + +torch = pytest.importorskip("torch") + + +def _load_paper_protocol_module(): + module_path = ( + Path(__file__).resolve().parents[2] + / "examples" + / "graph_tokenizer" + / "paper_protocol.py" + ) + spec = importlib.util.spec_from_file_location("paper_protocol_under_test", module_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +QM9_SPEC = SimpleNamespace( + canonical_name="qm9", + task_type="regression", + output_dim=16, + label_keys=( + "mu", "alpha", "homo", "lumo", "gap", "r2", "zpve", "u0", + "u298", "h298", "g298", "cv", "u0_atom", "u298_atom", + "h298_atom", "g298_atom", + ), +) + + +def test_paper_protocol_loads_native_tlx_model_classes(): + protocol = _load_paper_protocol_module() + + GraphBERT, GraphGTE = protocol._load_paper_model_classes() + + assert GraphBERT.__name__ == "GraphBERT" + assert GraphGTE.__name__ == "GraphGTE" + assert GraphBERT.__module__ == "gammagl.models.graph_bert" + assert GraphGTE.__module__ == "gammagl.models.graph_gte" + + +def test_paper_runtime_rejects_non_torch_tlx_backend(monkeypatch): + protocol = _load_paper_protocol_module() + monkeypatch.setenv("TL_BACKEND", "tensorflow") + + with pytest.raises(RuntimeError, match="TL_BACKEND=torch"): + protocol.require_paper_runtime(require_cuda=False) + + +@pytest.mark.parametrize("encoder_type", ["bert", "gte"]) +def test_paper_protocol_constructs_native_tlx_models(encoder_type): + protocol = _load_paper_protocol_module() + model = protocol._create_paper_model( + encoder_type=encoder_type, + vocab_size=24, + pad_token_id=0, + task_type="regression", + output_dim=2, + pooling="mean", + strict_architecture=False, + model_config={ + "hidden_size": 16, + "num_hidden_layers": 1, + "num_attention_heads": 4, + "intermediate_size": 32, + "max_position_embeddings": 16, + "dropout_rate": 0.0, + }, + allow_random_gte_init=(encoder_type == "gte"), + ) + input_ids = torch.tensor([[3, 5, 4, 0]], dtype=torch.long) + attention_mask = torch.tensor([[1, 1, 1, 0]], dtype=torch.long) + + assert model(input_ids, attention_mask, task="mlm").shape == (1, 4, 24) + assert model(input_ids, attention_mask, task="supervised").shape == (1, 2) + assert model.manifest()["encoder_type"] == encoder_type + + +def test_strict_paper_bert_ignores_caller_architecture_overrides(): + protocol = _load_paper_protocol_module() + model = protocol._create_paper_model( + encoder_type="bert", vocab_size=97, pad_token_id=0, + task_type="regression", output_dim=2, strict_architecture=True, + model_config={ + "hidden_size": 16, "num_hidden_layers": 1, + "num_attention_heads": 1, "intermediate_size": 32, + "max_position_embeddings": 16, + }) + manifest = model.manifest() + assert manifest["hidden_size"] == 512 + assert manifest["num_hidden_layers"] == 4 + assert manifest["num_attention_heads"] == 4 + assert manifest["intermediate_size"] == 2048 + assert manifest["max_position_embeddings"] == 8096 + assert model.token_embeddings.embeddings.shape == (97, 512) + + +def test_paper_gte_does_not_fallback_when_pretrained_loading_fails(monkeypatch): + protocol = _load_paper_protocol_module() + _, GraphGTE = protocol._load_paper_model_classes() + + def fail(cls, **kwargs): + raise RuntimeError("checksum mismatch") + + monkeypatch.setattr(GraphGTE, "from_pretrained", classmethod(fail)) + with pytest.raises(RuntimeError, match="checksum mismatch"): + protocol._create_paper_model( + encoder_type="gte", vocab_size=24, pad_token_id=1, + task_type="regression", output_dim=2) + + +def test_explicit_random_paper_gte_is_not_a_reproduction(): + protocol = _load_paper_protocol_module() + model = protocol._create_paper_model( + encoder_type="gte", vocab_size=24, pad_token_id=0, + task_type="regression", output_dim=2, strict_architecture=False, + allow_random_gte_init=True, + model_config={"hidden_size": 16, "num_hidden_layers": 1, + "num_attention_heads": 4, "intermediate_size": 32, + "max_position_embeddings": 16}) + assert model.manifest()["pretrained"] is False + assert model.manifest()["reproduction"] is False + + +def test_paper_protocol_rejects_non_fp32_model_parameters(): + protocol = _load_paper_protocol_module() + model = torch.nn.Linear(4, 2).half() + + with pytest.raises(RuntimeError, match="FP32 model parameters"): + protocol._require_fp32_parameters(torch, model) + + +def test_paper_protocol_accepts_qm9_multi_target_without_target_property(): + protocol = _load_paper_protocol_module() + args = SimpleNamespace(dataset="qm9", target_property=None, runs=5, seed=42) + + protocol.validate_paper_args(args, QM9_SPEC) + + +def test_paper_protocol_rejects_nonofficial_bert_position_limit(): + protocol = _load_paper_protocol_module() + args = SimpleNamespace( + dataset="qm9", target_property=None, runs=5, seed=42, + model="bert", max_position_embeddings=768) + + with pytest.raises(ValueError, match="max-position-embeddings=8096"): + protocol.validate_paper_args(args, QM9_SPEC) + + +def test_paper_protocol_rejects_legacy_qm9_single_target_selection(): + protocol = _load_paper_protocol_module() + args = SimpleNamespace(dataset="qm9", target_property="homo", runs=5, seed=42) + + with pytest.raises(ValueError, match="16-target"): + protocol.validate_paper_args(args, QM9_SPEC) + + +def test_paper_runtime_versions_accept_exact_paper_stack(): + protocol = _load_paper_protocol_module() + versions = { + "torch": "2.1.2+cu121", + "cuda": "12.1", + "dgl": "2.4.0+cu121", + "torch_geometric": "2.4.0", + } + + protocol.validate_paper_runtime_versions(versions) + + +@pytest.mark.parametrize( + ("field", "wrong_value"), + [ + ("torch", "2.11.0+cu128"), + ("cuda", "12.8"), + ("dgl", "1.1.3"), + ("torch_geometric", "2.8.0"), + ], +) +def test_paper_runtime_versions_reject_nonpaper_stack(field, wrong_value): + protocol = _load_paper_protocol_module() + versions = { + "torch": "2.1.2+cu121", + "cuda": "12.1", + "dgl": "2.4.0+cu121", + "torch_geometric": "2.4.0", + } + versions[field] = wrong_value + + with pytest.raises(RuntimeError, match=field): + protocol.validate_paper_runtime_versions(versions) + + +def test_scaler_fit_uses_train_only(): + protocol = _load_paper_protocol_module() + spec = SimpleNamespace( + canonical_name="qm9", + output_dim=2, + label_keys=("first", "second"), + ) + encoded = { + "train": {"labels": [[1.0, 10.0], [3.0, 14.0]]}, + "val": {"labels": [[5.0, 18.0]]}, + "test": {"labels": [[0.0, 8.0]]}, + } + + normalizer, output_dim = protocol._prepare_labels(encoded, spec, None) + + assert output_dim == 2 + assert normalizer == { + "mean": [2.0, 12.0], + "std": [1.0, 2.0], + "target_properties": ["first", "second"], + } + assert encoded["train"]["labels"] == [[-1.0, -1.0], [1.0, 1.0]] + assert encoded["val"]["labels"] == [[3.0, 3.0]] + assert encoded["test"]["labels"] == [[-2.0, -2.0]] + semantics = protocol._metric_semantics(spec, normalizer) + assert semantics["loss_space"] == "standardized" + assert semantics["metric_space"] == "raw" + assert semantics["target_normalization"]["type"] == "train_split_zscore" + + +def test_qm9_denormalization_restores_each_target_on_tensor_device(): + protocol = _load_paper_protocol_module() + normalized = torch.tensor([[-1.0, -1.0], [1.0, 1.0]]) + normalizer = { + "mean": [2.0, 12.0], + "std": [1.0, 2.0], + "target_properties": ["first", "second"], + } + + restored = protocol._denormalize_labels(torch, normalized, normalizer) + + assert restored.device == normalized.device + assert restored.dtype == normalized.dtype + assert torch.equal(restored, torch.tensor([[1.0, 10.0], [3.0, 14.0]])) + + +@pytest.mark.parametrize( + ("canonical_name", "expected"), + [("molhiv", ("logits", "probability", "none")), + ("peptides-struct", ("raw", "raw", "none"))], +) +def test_non_qm9_metric_spaces_are_dataset_specific(canonical_name, expected): + protocol = _load_paper_protocol_module() + semantics = protocol._metric_semantics( + SimpleNamespace(canonical_name=canonical_name), normalizer=None) + assert (semantics["loss_space"], semantics["metric_space"], + semantics["target_normalization"]["type"]) == expected + + +def test_qm9_prepare_labels_rejects_incorrect_target_width(): + protocol = _load_paper_protocol_module() + spec = SimpleNamespace( + canonical_name="qm9", + output_dim=2, + label_keys=("first", "second"), + ) + encoded = { + "train": {"labels": [[1.0, 10.0], [3.0, 14.0]]}, + "val": {"labels": [[5.0]]}, + "test": {"labels": [[0.0, 8.0]]}, + } + + with pytest.raises(ValueError, match="exactly 2 targets"): + protocol._prepare_labels(encoded, spec, None) + + +def test_qm9_prepare_labels_rejects_nonfinite_targets(): + protocol = _load_paper_protocol_module() + spec = SimpleNamespace( + canonical_name="qm9", + output_dim=2, + label_keys=("first", "second"), + ) + encoded = { + "train": {"labels": [[1.0, 10.0], [3.0, float("nan")]]}, + "val": {"labels": [[5.0, 18.0]]}, + "test": {"labels": [[0.0, 8.0]]}, + } + + with pytest.raises(ValueError, match="non-finite target"): + protocol._prepare_labels(encoded, spec, None) + + +def test_paper_protocol_uses_five_seeds_from_42_when_runs_is_unspecified(): + protocol = _load_paper_protocol_module() + args = SimpleNamespace(dataset="qm9", target_property=None, runs=None, seed=None) + + protocol.validate_paper_args(args, QM9_SPEC) + + assert protocol.paper_run_seeds(args) == [42, 43, 44, 45, 46] + + +def test_paper_training_options_are_read_from_argparse_namespace(): + protocol = _load_paper_protocol_module() + args = SimpleNamespace( + model="gte", + serialization="feuler", + pretrain_epoch=200, + n_epoch=200, + pretrain_lr=5e-5, + lr=1e-5, + batch_size=32, + weight_decay=0.1, + pretrain_warmup_ratio=0.12, + finetune_warmup_ratio=0.025, + pretrain_max_grad_norm=2.0, + finetune_max_grad_norm=0.5, + mask_prob=0.09, + patience=20, + max_position_embeddings=8192, + paper_amp="bf16", + paper_tf32=True, + paper_loss="huber", + molhiv_pos_weight=True, + ) + + options = protocol.paper_training_options(args) + + assert options == { + "encoder": "gte", + "serialization": "feuler", + "pretrain_epochs": 200, + "finetune_epochs": 200, + "pretrain_lr": 5e-5, + "finetune_lr": 1e-5, + "batch_size": 32, + "weight_decay": 0.1, + "pretrain_warmup_ratio": 0.12, + "finetune_warmup_ratio": 0.025, + "pretrain_max_grad_norm": 2.0, + "finetune_max_grad_norm": 0.5, + "mask_prob": 0.09, + "patience": 20, + "max_position_embeddings": 8192, + "amp_dtype": "bf16", + "allow_tf32": True, + "training_loss": "huber", + "molhiv_pos_weight": True, + } + + +def test_validation_not_used_for_training(): + protocol = _load_paper_protocol_module() + history = [ + {"epoch": 1, "val_metric": 0.4}, + {"epoch": 2, "val_metric": 0.2}, + {"epoch": 3, "val_metric": 0.3}, + ] + + assert protocol.select_best_epoch(history, higher_is_better=False) == 2 + assert "test" not in history[0] + assert "test" not in history[1] + assert "test" not in history[2] + + +def test_paper_regression_loss_backpropagates_from_half_precision_logits(): + protocol = _load_paper_protocol_module() + logits = torch.tensor([[1.0]], dtype=torch.float16, requires_grad=True) + labels = torch.tensor([[0.0]], dtype=torch.float32) + + loss = protocol._loss(torch, logits, labels, QM9_SPEC) + loss.backward() + + assert logits.grad is not None + + +def test_optional_regression_losses_leave_paper_default_unchanged(): + protocol = _load_paper_protocol_module() + spec = SimpleNamespace(canonical_name="peptides-struct") + logits = torch.tensor([[3.0, 2.0]]) + labels = torch.tensor([[0.0, float("nan")]]) + + default = protocol._loss(torch, logits, labels, spec) + l1 = protocol._loss(torch, logits, labels, spec, loss_name="l1") + huber = protocol._loss(torch, logits, labels, spec, loss_name="huber") + + assert default.item() == 9.0 + assert l1.item() == 3.0 + assert huber.item() == 2.5 + + +def test_optional_molhiv_positive_weight_changes_positive_example_loss(): + protocol = _load_paper_protocol_module() + spec = SimpleNamespace(canonical_name="molhiv") + logits = torch.tensor([[0.0]]) + labels = torch.tensor([[1.0]]) + + unweighted = protocol._loss(torch, logits, labels, spec) + weighted = protocol._loss( + torch, logits, labels, spec, pos_weight=torch.tensor(3.0)) + + assert weighted.item() == pytest.approx(unweighted.item() * 3.0) + + +def test_precision_options_reject_cuda_amp_on_cpu(): + protocol = _load_paper_protocol_module() + + assert protocol._resolve_precision_options( + torch, torch.device("cpu"), {"amp_dtype": "off"}) == { + "amp_dtype": "off", + "torch_dtype": None, + "grad_scaler": None, + "allow_tf32": False, + } + with pytest.raises(ValueError, match="CUDA device"): + protocol._resolve_precision_options( + torch, torch.device("cpu"), {"amp_dtype": "fp16"}) + + +def test_epoch_runtime_metrics_report_cpu_throughput(): + protocol = _load_paper_protocol_module() + + metrics = protocol._epoch_runtime_metrics( + torch, torch.device("cpu"), num_examples=10, seconds=2.0) + + assert metrics == { + "seconds": 2.0, + "examples_per_second": 5.0, + } + + +def test_paper_metric_details_report_each_qm9_target_mae(): + protocol = _load_paper_protocol_module() + spec = SimpleNamespace( + canonical_name="qm9", + label_keys=("first", "second"), + ) + + details = protocol.compute_paper_metric_details( + spec, + [[1.0, 4.0], [3.0, 8.0]], + [[2.0, 2.0], [1.0, 9.0]], + ) + + assert details == { + "metric": 1.5, + "per_target_mae": {"first": 1.5, "second": 1.5}, + } + + +def test_paper_mlm_mask_forces_one_token_when_random_selection_is_empty(monkeypatch): + protocol = _load_paper_protocol_module() + tokenizer = SimpleNamespace(special_tokens=SimpleNamespace( + pad_token_id=0, + cls_token_id=1, + sep_token_id=2, + component_sep_token_id=3, + mask_token_id=4, + )) + input_ids = torch.tensor([[1, 5, 6, 3, 7, 2, 0]]) + attention_mask = torch.tensor([[1, 1, 1, 1, 1, 1, 0]]) + monkeypatch.setattr( + torch, + "rand_like", + lambda tensor, dtype: torch.ones_like(tensor, dtype=dtype), + ) + + masked_ids, labels = protocol._mask_for_mlm( + torch, input_ids, attention_mask, tokenizer, mask_prob=0.09) + + selected = labels.ne(-100) + maskable = torch.tensor([[False, True, True, False, True, False, False]]) + assert selected.sum().item() == 1 + assert torch.all(selected <= maskable) + assert torch.equal(labels[selected], input_ids[selected]) + assert torch.all(masked_ids[selected].eq(tokenizer.special_tokens.mask_token_id)) + assert torch.equal(masked_ids[~selected], input_ids[~selected]) + + +def test_mlm_training_switches_tensorlayerx_model_to_train_mode(): + protocol = _load_paper_protocol_module() + + class TLXModeModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor(0.0)) + self.is_train = False + + def set_train(self): + self.is_train = True + + def forward(self, input_ids, attention_mask, task): + return self.weight.expand(len(input_ids), input_ids.shape[1], 8) + + tokenizer = SimpleNamespace(special_tokens=SimpleNamespace( + pad_token_id=0, + cls_token_id=1, + sep_token_id=2, + component_sep_token_id=3, + mask_token_id=4, + )) + model = TLXModeModel() + loader = [( + torch.tensor([[1, 5, 2]]), + torch.tensor([[1, 1, 1]]), + torch.tensor([[0.0]]), + )] + protocol._train_mlm( + torch, + model, + loader, + torch.optim.SGD(model.parameters(), lr=0.0), + SimpleNamespace(step=lambda: None), + tokenizer, + torch.device("cpu"), + max_grad_norm=10.0, + mask_prob=0.09, + ) + + assert model.is_train is True + + +def test_paper_evaluation_rejects_nonfinite_predictions(): + protocol = _load_paper_protocol_module() + + class NonFiniteModel(torch.nn.Module): + def forward(self, input_ids, attention_mask, task): + return torch.full((len(input_ids), 1), float("nan")) + + loader = [(torch.tensor([[1]]), torch.tensor([[1]]), torch.tensor([[0.0]]))] + + with pytest.raises(FloatingPointError, match="supervised logits"): + protocol._evaluate(torch, NonFiniteModel(), loader, QM9_SPEC, "cpu", None) + + +def test_paper_evaluation_switches_tensorlayerx_model_to_eval_mode(): + """Evaluation must disable TLX Dropout, which reads ``is_train``.""" + protocol = _load_paper_protocol_module() + + class TLXModeModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.is_train = True + + def set_eval(self): + self.is_train = False + + def forward(self, input_ids, attention_mask, task): + return torch.zeros((len(input_ids), 1), dtype=torch.float32) + + model = TLXModeModel() + loader = [(torch.tensor([[1]]), torch.tensor([[1]]), torch.tensor([[0.0]]))] + + protocol._evaluate( + torch, + model, + loader, + SimpleNamespace(canonical_name="qm9", label_keys=("target",)), + torch.device("cpu"), + normalizer=None, + ) + + assert model.is_train is False + + +def test_amp_evaluation_casts_predictions_to_fp32_before_denormalizing( + monkeypatch): + protocol = _load_paper_protocol_module() + + class HalfModel(torch.nn.Module): + def forward(self, input_ids, attention_mask, task): + return torch.full( + (len(input_ids), 1), 0.3333, dtype=torch.float16) + + loader = [( + torch.tensor([[1]]), + torch.tensor([[1]]), + torch.tensor([[0.0]], dtype=torch.float32), + )] + dtypes = [] + + def record_dtype(_torch, values, _normalizer): + dtypes.append(values.dtype) + return values + + monkeypatch.setattr(protocol, "_denormalize_labels", record_dtype) + protocol._evaluate( + torch, + HalfModel(), + loader, + SimpleNamespace(canonical_name="qm9", label_keys=("target",)), + torch.device("cpu"), + normalizer={"mean": [0.0], "std": [1.0]}, + ) + + assert dtypes == [torch.float32, torch.float32] + + +def test_qm9_evaluation_uses_standardized_mae_and_keeps_raw_diagnostics(): + protocol = _load_paper_protocol_module() + + class ZeroModel(torch.nn.Module): + def forward(self, input_ids, attention_mask, task): + return torch.zeros((len(input_ids), 2), dtype=torch.float32) + + result = protocol._evaluate( + torch, + ZeroModel(), + [(torch.tensor([[1]]), torch.tensor([[1]]), + torch.tensor([[1.0, 1.0]]))], + SimpleNamespace( + canonical_name="qm9", label_keys=("first", "second")), + torch.device("cpu"), + normalizer={"mean": [100.0, 1000.0], "std": [10.0, 100.0]}, + ) + + assert result["metric"] == 55.0 + assert result["metric_space"] == "raw" + assert result["per_target_mae"] == {"first": 1.0, "second": 1.0} + assert result["per_target_mae_raw"] == {"first": 10.0, "second": 100.0} + + +def test_paper_encoding_keeps_ragged_sequences_until_batch_collation(): + protocol = _load_paper_protocol_module() + + class Tokenizer: + special_tokens = SimpleNamespace(pad_token_id=0) + + @staticmethod + def encode_graph(graph): + return SimpleNamespace(input_ids=list(graph.tokens)) + + def graph(tokens, label): + return SimpleNamespace(tokens=tokens, y=[label]) + + splits = { + "train": [graph([3, 8, 4], 1.0), graph([3, 9, 10, 11, 4], 2.0)], + "val": [graph([3, 4], 3.0)], + "test": [graph([3, 12, 4], 4.0)], + } + + encoded, maximum = protocol._encode_splits( + Tokenizer(), splits, max_position_embeddings=16) + + assert maximum == 5 + assert encoded["train"]["input_ids"] == [ + [3, 8, 4], + [3, 9, 10, 11, 4], + ] + assert "attention_mask" not in encoded["train"] + + +def test_paper_encoded_splits_reuse_local_cache_without_reserializing(tmp_path): + protocol = _load_paper_protocol_module() + + class CountingTokenizer: + calls = 0 + + @classmethod + def encode_graph(cls, graph): + cls.calls += 1 + return SimpleNamespace(input_ids=list(graph.tokens)) + + def graph(tokens, label): + return SimpleNamespace(tokens=tokens, y=[label]) + + splits = { + "train": [graph([3, 8, 4], 1.0)], + "val": [graph([3, 9, 4], 2.0)], + "test": [graph([3, 10, 4], 3.0)], + } + cache_path = tmp_path / "encoded.pkl" + + first, first_max = protocol._load_or_encode_splits( + CountingTokenizer(), splits, 16, cache_path, cache_key="fixture-v1") + second, second_max = protocol._load_or_encode_splits( + CountingTokenizer(), splits, 16, cache_path, cache_key="fixture-v1") + + assert CountingTokenizer.calls == 3 + assert second == first + assert second_max == first_max == 3 + + +def test_encoded_cache_key_changes_when_validation_content_changes(): + protocol = _load_paper_protocol_module() + tokenizer = SimpleNamespace( + _cache_key="train-v1", + serializer=SimpleNamespace(name="feuler", frequency_map={}), + bpe=SimpleNamespace(codebook=SimpleNamespace( + merge_rules=[], vocab_size=8)), + ) + + def graph(label): + return SimpleNamespace( + edge_index=[[0], [1]], + x=[13, 17], + edge_attr=[2], + num_nodes=2, + y=[label], + ) + + first_splits = { + "train": [graph(0.0)], + "val": [graph(1.0)], + "test": [graph(2.0)], + } + second_splits = { + **first_splits, + "val": [graph(99.0)], + } + + first = protocol._encoded_splits_cache_key(tokenizer, first_splits, 16) + second = protocol._encoded_splits_cache_key(tokenizer, second_splits, 16) + + assert first != second + + +def test_paper_encoding_chunks_native_batch_calls(): + protocol = _load_paper_protocol_module() + + class Tokenizer: + calls = [] + + @classmethod + def batch_encode_graphs(cls, graphs): + cls.calls.append(len(graphs)) + return [SimpleNamespace(input_ids=list(graph.tokens)) for graph in graphs] + + graphs = [ + SimpleNamespace(tokens=[3, index, 4], y=[float(index)]) + for index in range(5) + ] + splits = {"train": graphs, "val": graphs[:1], "test": graphs[:1]} + + encoded, maximum = protocol._encode_splits( + Tokenizer(), splits, 16, encoding_batch_size=2) + + assert Tokenizer.calls == [2, 2, 1, 1, 1] + assert len(encoded["train"]["input_ids"]) == 5 + assert maximum == 3 + + +def test_paper_loader_pads_only_to_each_batch_maximum(): + protocol = _load_paper_protocol_module() + split = { + "input_ids": [[3, 8, 4], [3, 9, 10, 11, 4], [3, 4]], + "labels": [[1.0], [2.0], [3.0]], + } + + loader = protocol._make_loader( + torch, + split, + batch_size=2, + shuffle=False, + pad_token_id=0, + num_workers=0, + pin_memory=False, + ) + first_ids, first_mask, first_labels = next(iter(loader)) + batches = list(loader) + second_ids, second_mask, second_labels = batches[1] + + assert first_ids.tolist() == [[3, 8, 4, 0, 0], [3, 9, 10, 11, 4]] + assert first_mask.tolist() == [[1, 1, 1, 0, 0], [1, 1, 1, 1, 1]] + assert first_labels.tolist() == [[1.0], [2.0]] + assert second_ids.tolist() == [[3, 4]] + assert second_mask.tolist() == [[1, 1]] + assert second_labels.tolist() == [[3.0]] + + +def test_dataloader_worker_seed_does_not_consume_model_rng(): + protocol = _load_paper_protocol_module() + split = { + "input_ids": [[3, 4]], + "labels": [[1.0]], + } + torch.manual_seed(123) + expected = torch.rand(1) + torch.manual_seed(123) + loader = protocol._make_loader( + torch, + split, + batch_size=1, + shuffle=False, + num_workers=0, + bucket_by_length=False, + ) + + next(iter(loader)) + actual = torch.rand(1) + + assert torch.equal(actual, expected) + + +def test_paper_memory_plan_keeps_effective_batch_for_short_sequences(): + protocol = _load_paper_protocol_module() + + plan = protocol._resolve_memory_plan( + encoder_type="gte", + effective_batch_size=16, + sequence_length=165, + available_bytes=24 * 1024 ** 3, + ) + + assert plan["micro_batch_size"] == 16 + assert plan["gradient_accumulation_steps"] == 1 + assert plan["effective_batch_size"] == 16 + assert plan["estimated_peak_bytes"] < plan["budget_bytes"] + + +def test_cuda_memory_query_uses_current_index_for_bare_cuda_device(monkeypatch): + protocol = _load_paper_protocol_module() + monkeypatch.setattr(torch.cuda, "current_device", lambda: 3) + + assert protocol._cuda_memory_query_device(torch, torch.device("cuda")) == 3 + assert ( + protocol._cuda_memory_query_device(torch, torch.device("cuda:2")) + == torch.device("cuda:2") + ) + + +def test_paper_memory_plan_uses_gradient_accumulation_when_needed(): + protocol = _load_paper_protocol_module() + + plan = protocol._resolve_memory_plan( + encoder_type="bert", + effective_batch_size=32, + sequence_length=768, + available_bytes=6 * 1024 ** 3, + ) + + assert 1 <= plan["micro_batch_size"] < 32 + assert 32 % plan["micro_batch_size"] == 0 + assert ( + plan["micro_batch_size"] + * plan["gradient_accumulation_steps"] + == 32 + ) + + +def test_paper_memory_plan_accounts_for_mixed_precision_activations(): + protocol = _load_paper_protocol_module() + + fp32 = protocol._estimate_training_peak_bytes( + "bert", micro_batch_size=8, sequence_length=256, + activation_bytes_per_element=4) + mixed = protocol._estimate_training_peak_bytes( + "bert", micro_batch_size=8, sequence_length=256, + activation_bytes_per_element=2) + + assert mixed < fp32 + + +def test_resume_skips_finetuning_when_saved_state_already_hit_patience(): + protocol = _load_paper_protocol_module() + + assert protocol._should_skip_finetuning( + resume_phase="finetune", stale_epochs=20, patience=20) + assert protocol._should_skip_finetuning( + resume_phase="finetune_complete", stale_epochs=0, patience=20) + assert not protocol._should_skip_finetuning( + resume_phase="finetune", stale_epochs=19, patience=20) + + +def test_paper_memory_plan_rejects_impossible_dense_attention(): + protocol = _load_paper_protocol_module() + + with pytest.raises(RuntimeError, match="dense attention"): + protocol._resolve_memory_plan( + encoder_type="gte", + effective_batch_size=16, + sequence_length=8192, + available_bytes=24 * 1024 ** 3, + ) + + +def test_supervised_training_accumulates_to_the_effective_batch(): + protocol = _load_paper_protocol_module() + + class TinyModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor([[0.5]])) + + def forward(self, input_ids, attention_mask, task): + assert task == "supervised" + return input_ids.float()[:, :1] @ self.weight + + model = TinyModel() + loader = [ + (torch.tensor([[value]]), torch.ones((1, 1), dtype=torch.long), + torch.tensor([[0.0]])) + for value in (1, 2, 3, 4) + ] + optimizer = torch.optim.SGD(model.parameters(), lr=0.01) + step_calls = [] + original_step = optimizer.step + + def counted_step(*args, **kwargs): + step_calls.append(True) + return original_step(*args, **kwargs) + + optimizer.step = counted_step + scheduler = SimpleNamespace(step=lambda: None) + + protocol._train_supervised( + torch, + model, + loader, + optimizer, + scheduler, + QM9_SPEC, + "cpu", + max_grad_norm=10.0, + gradient_accumulation_steps=2, + ) + + assert len(step_calls) == 2 + + +def test_supervised_training_uses_explicit_regression_loss(): + protocol = _load_paper_protocol_module() + + class TinyModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor([[3.0]])) + + def forward(self, input_ids, attention_mask, task): + return input_ids.float()[:, :1] @ self.weight + + model = TinyModel() + loader = [( + torch.tensor([[1]]), + torch.ones((1, 1), dtype=torch.long), + torch.tensor([[0.0]]), + )] + optimizer = torch.optim.SGD(model.parameters(), lr=0.0) + + loss = protocol._train_supervised( + torch, + model, + loader, + optimizer, + SimpleNamespace(step=lambda: None), + SimpleNamespace(canonical_name="qm9"), + torch.device("cpu"), + max_grad_norm=10.0, + loss_name="l1", + ) + + assert loss == 3.0 + + +def test_supervised_training_switches_tensorlayerx_model_to_train_mode(): + protocol = _load_paper_protocol_module() + + class TLXModeModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor([[1.0]])) + self.is_train = False + + def set_train(self): + self.is_train = True + + def forward(self, input_ids, attention_mask, task): + return input_ids.float()[:, :1] @ self.weight + + model = TLXModeModel() + loader = [( + torch.tensor([[1]]), torch.ones((1, 1), dtype=torch.long), + torch.tensor([[0.0]]), + )] + protocol._train_supervised( + torch, + model, + loader, + torch.optim.SGD(model.parameters(), lr=0.0), + SimpleNamespace(step=lambda: None), + SimpleNamespace(canonical_name="qm9"), + torch.device("cpu"), + max_grad_norm=10.0, + ) + + assert model.is_train is True + + +def test_gradient_accumulation_weights_uneven_microbatches_by_examples(): + protocol = _load_paper_protocol_module() + + class TinyModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor([[1.0]])) + + def forward(self, input_ids, attention_mask, task): + return input_ids.float()[:, :1] @ self.weight + + model = TinyModel() + loader = [ + ( + torch.tensor([[1], [1]]), + torch.ones((2, 1), dtype=torch.long), + torch.zeros((2, 1)), + ), + ( + torch.tensor([[3]]), + torch.ones((1, 1), dtype=torch.long), + torch.zeros((1, 1)), + ), + ] + optimizer = torch.optim.SGD(model.parameters(), lr=0.0) + gradients = [] + original_step = optimizer.step + + def record_gradient(*args, **kwargs): + gradients.append(float(model.weight.grad.item())) + return original_step(*args, **kwargs) + + optimizer.step = record_gradient + + protocol._train_supervised( + torch, + model, + loader, + optimizer, + SimpleNamespace(step=lambda: None), + SimpleNamespace(canonical_name="qm9"), + torch.device("cpu"), + max_grad_norm=100.0, + gradient_accumulation_steps=2, + ) + + assert gradients == pytest.approx([22.0 / 3.0]) + + +def test_fp16_grad_scaler_is_new_for_each_independent_run(): + protocol = _load_paper_protocol_module() + + class FakeGradScaler: + pass + + fake_torch = SimpleNamespace( + cuda=SimpleNamespace( + amp=SimpleNamespace(GradScaler=lambda enabled: FakeGradScaler()))) + + first = protocol._new_grad_scaler(fake_torch, "fp16") + second = protocol._new_grad_scaler(fake_torch, "fp16") + + assert isinstance(first, FakeGradScaler) + assert isinstance(second, FakeGradScaler) + assert first is not second + assert protocol._new_grad_scaler(fake_torch, "bf16") is None + + +def test_fp16_overflow_skips_optimizer_and_scheduler_step(): + protocol = _load_paper_protocol_module() + + class TinyModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.tensor([[1.0]])) + + def forward(self, input_ids, attention_mask, task): + return input_ids.float()[:, :1] @ self.weight + + model = TinyModel() + loader = [( + torch.tensor([[1]]), + torch.ones((1, 1), dtype=torch.long), + torch.zeros((1, 1)), + )] + optimizer = torch.optim.SGD(model.parameters(), lr=0.1) + optimizer_steps = [] + scheduler_steps = [] + optimizer.step = lambda *_args, **_kwargs: optimizer_steps.append(True) + + class OverflowScaler: + scale_value = 8.0 + + @staticmethod + def scale(loss): + return loss + + @staticmethod + def unscale_(_optimizer): + model.weight.grad.fill_(float("inf")) + + @staticmethod + def step(_optimizer): + return None + + def update(self): + self.scale_value /= 2 + + def get_scale(self): + return self.scale_value + + protocol._train_supervised( + torch, + model, + loader, + optimizer, + SimpleNamespace(step=lambda: scheduler_steps.append(True)), + SimpleNamespace(canonical_name="qm9"), + torch.device("cpu"), + max_grad_norm=1.0, + grad_scaler=OverflowScaler(), + ) + + assert optimizer_steps == [] + assert scheduler_steps == [] + + +def test_molhiv_positive_weight_uses_training_labels_only(): + protocol = _load_paper_protocol_module() + encoded_train = { + "labels": [[0.0], [0.0], [0.0], [1.0]], + } + + weight = protocol._molhiv_positive_weight(torch, encoded_train, "cpu") + + assert weight.item() == 3.0 + + +def test_best_checkpoint_is_lightweight_and_restores_model(tmp_path): + protocol = _load_paper_protocol_module() + model = torch.nn.Linear(2, 1) + path = tmp_path / "best.pt" + + protocol.save_paper_checkpoint( + path, + torch, + model, + epoch=7, + best_metric=0.25, + normalizer={"mean": [1.0], "std": [2.0]}, + ) + state = protocol._torch_load(torch, path, device="cpu") + + assert set(state) == { + "checkpoint_kind", + "model", + "epoch", + "best_metric", + "normalizer", + } + assert state["checkpoint_kind"] == "best_model" + assert "optimizer" not in state + assert "scheduler" not in state + with torch.no_grad(): + model.weight.zero_() + restored = protocol.restore_paper_checkpoint( + path, torch, model, device="cpu") + assert restored["epoch"] == 7 + assert torch.count_nonzero(model.weight).item() > 0 + + +def test_checkpoint_rejects_a_different_experiment_fingerprint(tmp_path): + protocol = _load_paper_protocol_module() + model = torch.nn.Linear(1, 1) + path = tmp_path / "best.pt" + protocol.save_paper_checkpoint( + path, + torch, + model, + epoch=1, + best_metric=0.5, + normalizer=None, + experiment_fingerprint="experiment-a", + ) + + with pytest.raises(ValueError, match="fingerprint"): + protocol.restore_paper_checkpoint( + path, + torch, + model, + device="cpu", + expected_fingerprint="experiment-b", + ) + + +def test_resume_checkpoint_roundtrips_training_and_rng_state(tmp_path): + protocol = _load_paper_protocol_module() + model = torch.nn.Linear(2, 1) + optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3) + scheduler = torch.optim.lr_scheduler.LambdaLR( + optimizer, lambda _: 1.0) + loss = model(torch.ones((1, 2))).sum() + loss.backward() + optimizer.step() + scheduler.step() + path = tmp_path / "last_state.pt" + + protocol.save_paper_resume_state( + path, + torch, + model, + optimizer, + scheduler, + phase="pretrain", + epoch=3, + extra={"history": [{"epoch": 1}]}, + ) + state = protocol._torch_load(torch, path, device="cpu") + + assert state["checkpoint_kind"] == "resume_state" + assert state["phase"] == "pretrain" + assert state["epoch"] == 3 + assert "optimizer" in state + assert "scheduler" in state + assert "rng_state" in state + restored = protocol.restore_paper_resume_state( + path, torch, model, optimizer, scheduler, device="cpu") + assert restored["extra"]["history"] == [{"epoch": 1}] + + +def test_resume_checkpoint_roundtrips_amp_scaler_state(tmp_path): + protocol = _load_paper_protocol_module() + model = torch.nn.Linear(1, 1) + optimizer = torch.optim.SGD(model.parameters(), lr=0.1) + scheduler = torch.optim.lr_scheduler.LambdaLR( + optimizer, lambda _: 1.0) + + class FakeScaler: + def __init__(self, scale): + self.scale = scale + + def state_dict(self): + return {"scale": self.scale} + + def load_state_dict(self, state): + self.scale = state["scale"] + + path = tmp_path / "amp_state.pt" + protocol.save_paper_resume_state( + path, + torch, + model, + optimizer, + scheduler, + phase="pretrain", + epoch=2, + grad_scaler=FakeScaler(1024.0), + ) + restored_scaler = FakeScaler(1.0) + + protocol.restore_paper_resume_state( + path, + torch, + model, + optimizer, + scheduler, + device="cpu", + grad_scaler=restored_scaler, + ) + + assert restored_scaler.scale == 1024.0 + + +def test_pretrain_complete_state_restores_amp_scaler_before_finetuning(): + protocol = _load_paper_protocol_module() + + class FakeScaler: + scale = 1.0 + + def load_state_dict(self, state): + self.scale = state["scale"] + + scaler = FakeScaler() + protocol._restore_grad_scaler_state( + {"grad_scaler": {"scale": 512.0}}, scaler) + + assert scaler.scale == 512.0 + + +def test_run_fingerprint_changes_with_seed_and_run_index(): + protocol = _load_paper_protocol_module() + training_options = {"encoder": "gte", "batch_size": 32} + spec = SimpleNamespace( + canonical_name="qm9", + task_type="regression", + output_dim=16, + ) + + first = protocol._experiment_fingerprint( + training_options, spec, "encoded", 100, "mean", seed=42, run_index=0) + different_seed = protocol._experiment_fingerprint( + training_options, spec, "encoded", 100, "mean", seed=43, run_index=0) + different_run = protocol._experiment_fingerprint( + training_options, spec, "encoded", 100, "mean", seed=42, run_index=1) + + assert len({first, different_seed, different_run}) == 3 + + +def test_tiny_paper_protocol_runs_pretrain_finetune_and_five_runs( + tmp_path, monkeypatch): + protocol = _load_paper_protocol_module() + + class TinyModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.embedding = torch.nn.Embedding(32, 4) + self.task_head = torch.nn.Linear(4, 16) + self.mlm_head = torch.nn.Linear(4, 32) + + def forward(self, input_ids, attention_mask, task): + hidden = self.embedding(input_ids) + if task == "mlm": + return self.mlm_head(hidden) + mask = attention_mask.unsqueeze(-1).float() + pooled = (hidden * mask).sum(1) / mask.sum(1).clamp_min(1) + return self.task_head(pooled) + + @staticmethod + def manifest(): + return {"encoder_type": "bert", "tiny": True} + + class Tokenizer: + special_tokens = SimpleNamespace( + pad_token_id=0, + cls_token_id=1, + sep_token_id=2, + component_sep_token_id=3, + mask_token_id=4, + ) + serializer = SimpleNamespace(name="feuler", frequency_map={}) + bpe = SimpleNamespace(codebook=SimpleNamespace( + merge_rules=[], vocab_size=32)) + max_token_id = 31 + _cache_key = "tiny-tokenizer" + _cache_status = "miss" + _cache_path = str(tmp_path / "tokenizer.pkl") + + @staticmethod + def validate_model_vocab(vocab_size): + assert vocab_size == 32 + + @staticmethod + def encode_graph(graph): + return SimpleNamespace(input_ids=list(graph.tokens)) + + labels = [float(index) for index in range(16)] + + def graph(token, offset): + return SimpleNamespace( + tokens=[1, token, 2], + y=[value + offset for value in labels], + ) + + splits = { + "train": [graph(5, 0.0), graph(6, 1.0)], + "val": [graph(7, 0.5)], + "test": [graph(8, 0.25)], + } + splits["test"][0].tokens = [1, 8, 9, 10, 11, 12, 13, 14, 2] + args = SimpleNamespace( + dataset="qm9", + target_property=None, + runs=5, + seed=42, + model="bert", + pooling="mean", + device="cpu", + num_workers=0, + output_dir=str(tmp_path / "output"), + data_root=str(tmp_path / "data"), + paper_cache_root=str(tmp_path / "cache"), + resume=False, + ) + training_options = { + "encoder": "bert", + "serialization": "feuler", + "pretrain_epochs": 1, + "finetune_epochs": 3, + "pretrain_lr": 1e-4, + "finetune_lr": 1e-5, + "patience": 0, + "batch_size": 2, + "weight_decay": 0.1, + "pretrain_warmup_ratio": 0.12, + "finetune_warmup_ratio": 0.025, + "pretrain_max_grad_norm": 2.0, + "finetune_max_grad_norm": 0.5, + "mask_prob": 0.09, + "max_position_embeddings": 16, + } + monkeypatch.setattr(protocol, "_torch", lambda: torch) + monkeypatch.setattr( + protocol, "paper_training_options", lambda _args: training_options) + model_configs = [] + + def create_tiny_model(**kwargs): + model_configs.append(kwargs["model_config"]) + return TinyModel() + + monkeypatch.setattr(protocol, "_create_paper_model", create_tiny_model) + scaler_runs = [] + monkeypatch.setattr( + protocol, + "_new_grad_scaler", + lambda _torch, amp_dtype: scaler_runs.append(amp_dtype) or None, + ) + + def fit_train_only(_args, graphs): + assert graphs == splits["train"] + assert all(graph not in graphs for graph in splits["val"] + splits["test"]) + return Tokenizer() + + summary = protocol.run_paper_experiment( + args, QM9_SPEC, splits, fit_train_only) + + assert len(summary["runs"]) == 5 + assert model_configs == [{"max_position_embeddings": 16}] * 5 + assert scaler_runs == ["off"] * 5 + assert all(run["test_evaluations"] == 1 for run in summary["runs"]) + assert all(Path(run["checkpoint_path"]).is_file() for run in summary["runs"]) + assert summary["cache"]["encoded_status"] == "miss" + assert summary["precision"] == { + "amp_dtype": "off", + "allow_tf32": False, + } + assert summary["loss_space"] == "standardized" + assert summary["metric_space"] == "raw" + assert summary["target_normalization"]["type"] == "train_split_zscore" + + args.resume = True + + def fail_if_completed_run_retrains(*_args, **_kwargs): + raise AssertionError("a finetune_complete run must not train again") + + monkeypatch.setattr( + protocol, "_train_supervised", fail_if_completed_run_retrains) + resumed = protocol.run_paper_experiment( + args, QM9_SPEC, splits, fit_train_only) + + assert all(run["resumed"] for run in resumed["runs"]) diff --git a/tests/models/test_graph_transformer.py b/tests/models/test_graph_transformer.py new file mode 100644 index 000000000..b3029e7f0 --- /dev/null +++ b/tests/models/test_graph_transformer.py @@ -0,0 +1,1220 @@ +import importlib.util +import json +import pickle +import gzip +import sys +import tempfile +import types +from types import SimpleNamespace +from pathlib import Path + +import numpy as np +import pytest +import tensorlayerx as tlx + + +class _OptionalTorch: + """Load torch only in tests that exercise the torch-specific path.""" + + def __getattr__(self, name): + return getattr(pytest.importorskip("torch"), name) + + +torch = _OptionalTorch() + + +def _load_model_module(filename, module_name): + module_path = Path(__file__).resolve().parents[2] / "gammagl" / "models" / filename + spec = importlib.util.spec_from_file_location(module_name, module_path) + module = importlib.util.module_from_spec(spec) + sys.modules[module_name] = module + spec.loader.exec_module(module) + return module + + +def _load_trainer_module(): + module_path = Path(__file__).resolve().parents[2] / "examples" / "graph_tokenizer" / "graph_tokenizer_trainer.py" + spec = importlib.util.spec_from_file_location("graph_tokenizer_trainer_under_test", module_path) + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +def _qm9_properties(offset=0.0): + names = ( + "mu", "alpha", "homo", "lumo", "gap", "r2", "zpve", "u0", + "u298", "h298", "g298", "cv", "u0_atom", "u298_atom", + "h298_atom", "g298_atom", + ) + return {name: offset + (index + 1) / 10.0 for index, name in enumerate(names)} + + +def test_graph_transformer_public_interfaces_are_importable(): + GraphBERT = _load_model_module("graph_bert.py", "graph_bert_public_import_test").GraphBERT + _load_model_module("graph_bert.py", "graph_bert") + GraphGTE = _load_model_module("graph_gte.py", "graph_gte_public_import_test").GraphGTE + + bert = GraphBERT(vocab_size=32, output_dim=2) + gte = GraphGTE(vocab_size=32, output_dim=2) + + assert bert.vocab_size == 32 + assert bert.output_dim == 2 + assert gte.vocab_size == 32 + assert gte.output_dim == 2 + + +def test_graph_bert_forward_returns_task_and_mlm_logits(): + GraphBERT = _load_model_module("graph_bert.py", "graph_bert_under_test").GraphBERT + + model = GraphBERT( + vocab_size=32, + output_dim=3, + hidden_size=16, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=32, + max_position_embeddings=16, + ) + input_ids = tlx.convert_to_tensor([[3, 5, 6, 4], [3, 7, 0, 0]], dtype=tlx.int64) + attention_mask = tlx.convert_to_tensor([[1, 1, 1, 1], [1, 1, 0, 0]], dtype=tlx.int64) + + outputs = model(input_ids, attention_mask=attention_mask) + + assert tlx.get_tensor_shape(outputs["logits"]) == [2, 3] + assert tlx.get_tensor_shape(outputs["mlm_logits"]) == [2, 4, 32] + assert tlx.get_tensor_shape(outputs["last_hidden_state"]) == [2, 4, 16] + assert tlx.get_tensor_shape(outputs["pooled_output"]) == [2, 16] + + +def test_graph_bert_supports_paper_task_outputs(): + GraphBERT = _load_model_module( + "graph_bert.py", "graph_bert_paper_outputs_test").GraphBERT + model = GraphBERT( + vocab_size=32, + output_dim=3, + hidden_size=16, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=32, + max_position_embeddings=16, + ) + input_ids = tlx.convert_to_tensor([[3, 5, 6, 4]], dtype=tlx.int64) + attention_mask = tlx.convert_to_tensor([[1, 1, 1, 1]], dtype=tlx.int64) + + mlm_logits = model(input_ids, attention_mask=attention_mask, task="mlm") + task_logits = model(input_ids, attention_mask=attention_mask, task="supervised") + + assert tlx.get_tensor_shape(mlm_logits) == [1, 4, 32] + assert tlx.get_tensor_shape(task_logits) == [1, 3] + + +@pytest.mark.parametrize( + ("filename", "module_name", "class_name"), + [ + ("graph_bert.py", "graph_bert_task_routing_test", "GraphBERT"), + ("graph_gte.py", "graph_gte_task_routing_test", "GraphGTE"), + ], +) +def test_graph_transformers_only_compute_the_requested_output_head( + filename, module_name, class_name, monkeypatch): + _load_model_module("graph_bert.py", "graph_bert") + model_class = getattr(_load_model_module(filename, module_name), class_name) + model = model_class( + vocab_size=32, + output_dim=2, + hidden_size=16, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=32, + max_position_embeddings=16, + dropout_rate=0.0, + ) + input_ids = tlx.convert_to_tensor([[3, 5, 4]], dtype=tlx.int64) + attention_mask = tlx.convert_to_tensor([[1, 1, 1]], dtype=tlx.int64) + + def fail_mlm(*args, **kwargs): + raise AssertionError("MLM head must not run for supervised inference") + + monkeypatch.setattr(model.mlm_head, "forward", fail_mlm) + supervised = model( + input_ids, attention_mask=attention_mask, task="supervised") + assert tlx.get_tensor_shape(supervised) == [1, 2] + + monkeypatch.undo() + + def fail_task(*args, **kwargs): + raise AssertionError("task head must not run for MLM inference") + + monkeypatch.setattr(model, "_task_logits", fail_task) + mlm = model(input_ids, attention_mask=attention_mask, task="mlm") + assert tlx.get_tensor_shape(mlm) == [1, 3, 32] + + +def test_graph_bert_padding_tokens_do_not_change_supervised_output(): + GraphBERT = _load_model_module( + "graph_bert.py", "graph_bert_padding_test").GraphBERT + model = GraphBERT( + vocab_size=32, + output_dim=2, + hidden_size=16, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=32, + max_position_embeddings=16, + dropout_rate=0.0, + ) + model.set_eval() + attention_mask = tlx.convert_to_tensor([[1, 1, 1, 0]], dtype=tlx.int64) + first = tlx.convert_to_tensor([[3, 5, 4, 0]], dtype=tlx.int64) + second = tlx.convert_to_tensor([[3, 5, 4, 17]], dtype=tlx.int64) + + first_logits = model(first, attention_mask=attention_mask, task="supervised") + second_logits = model(second, attention_mask=attention_mask, task="supervised") + + assert torch.allclose(first_logits, second_logits, atol=1e-6, rtol=1e-6) + + +def test_graph_bert_manifest_reports_strict_paper_defaults(): + GraphBERT = _load_model_module( + "graph_bert.py", "graph_bert_manifest_test").GraphBERT + model = GraphBERT(vocab_size=2200, output_dim=16) + + manifest = model.manifest() + + assert manifest["encoder_type"] == "bert" + assert manifest["model_name"] == "bert-small" + assert manifest["hidden_size"] == 512 + assert manifest["num_hidden_layers"] == 4 + assert manifest["num_attention_heads"] == 4 + assert manifest["intermediate_size"] == 2048 + assert manifest["hidden_act"] == "gelu" + assert manifest["hidden_dropout_prob"] == 0.1 + assert manifest["attention_probs_dropout_prob"] == 0.1 + assert manifest["position_embedding_type"] == "absolute" + assert manifest["max_position_embeddings"] == 768 + assert manifest["layer_norm_eps"] == 1e-12 + assert 14_000_000 < manifest["total_parameters"] < 20_000_000 + assert manifest["trainable_parameters"] == manifest["total_parameters"] + + +def test_graph_bert_is_native_scratch_model_for_the_graph_bpe_vocabulary(): + GraphBERT = _load_model_module( + "graph_bert.py", "graph_bert_graph_vocabulary_test").GraphBERT + model = GraphBERT(vocab_size=37, output_dim=2) + + assert not hasattr(GraphBERT, "from_pretrained") + assert model.token_embeddings.embeddings.shape == (37, 512) + assert model.position_embeddings.embeddings.shape == (768, 512) + assert model.token_type_embeddings.embeddings.shape == (2, 512) + + +def test_graph_bert_matches_hf_bert_encoder_when_given_identical_weights(): + transformers = pytest.importorskip("transformers") + GraphBERT = _load_model_module( + "graph_bert.py", "graph_bert_hf_reference_test").GraphBERT + model = GraphBERT( + vocab_size=31, output_dim=2, hidden_size=16, num_hidden_layers=1, + num_attention_heads=4, intermediate_size=32, + max_position_embeddings=16, dropout_rate=0.0, + attention_dropout_rate=0.0) + reference = transformers.BertModel(transformers.BertConfig( + vocab_size=31, hidden_size=16, num_hidden_layers=1, + num_attention_heads=4, intermediate_size=32, + max_position_embeddings=16, hidden_dropout_prob=0.0, + attention_probs_dropout_prob=0.0, layer_norm_eps=1e-12)).eval() + + def assign(target, source): + value = tlx.convert_to_tensor(source.detach().numpy(), dtype=target.dtype) + if hasattr(target, "assign"): + target.assign(value) + else: + target.data.copy_(value) + + for target, source in ( + (model.token_embeddings.embeddings, reference.embeddings.word_embeddings.weight), + (model.position_embeddings.embeddings, reference.embeddings.position_embeddings.weight), + (model.token_type_embeddings.embeddings, reference.embeddings.token_type_embeddings.weight), + (model.embedding_norm.gamma, reference.embeddings.LayerNorm.weight), + (model.embedding_norm.beta, reference.embeddings.LayerNorm.bias)): + assign(target, source) + tlx_layer = model.encoder_layers[0] + hf_layer = reference.encoder.layer[0] + for target, source in ( + (tlx_layer.attention.q_proj.weights, hf_layer.attention.self.query.weight), + (tlx_layer.attention.q_proj.biases, hf_layer.attention.self.query.bias), + (tlx_layer.attention.k_proj.weights, hf_layer.attention.self.key.weight), + (tlx_layer.attention.k_proj.biases, hf_layer.attention.self.key.bias), + (tlx_layer.attention.v_proj.weights, hf_layer.attention.self.value.weight), + (tlx_layer.attention.v_proj.biases, hf_layer.attention.self.value.bias), + (tlx_layer.attention.out_proj.weights, hf_layer.attention.output.dense.weight), + (tlx_layer.attention.out_proj.biases, hf_layer.attention.output.dense.bias), + (tlx_layer.attention_norm.gamma, hf_layer.attention.output.LayerNorm.weight), + (tlx_layer.attention_norm.beta, hf_layer.attention.output.LayerNorm.bias), + (tlx_layer.ffn_in.weights, hf_layer.intermediate.dense.weight), + (tlx_layer.ffn_in.biases, hf_layer.intermediate.dense.bias), + (tlx_layer.ffn_out.weights, hf_layer.output.dense.weight), + (tlx_layer.ffn_out.biases, hf_layer.output.dense.bias), + (tlx_layer.ffn_norm.gamma, hf_layer.output.LayerNorm.weight), + (tlx_layer.ffn_norm.beta, hf_layer.output.LayerNorm.bias)): + assign(target, source) + + model.set_eval() + input_ids = torch.tensor([[2, 4, 5, 0]], dtype=torch.long) + attention_mask = torch.tensor([[1, 1, 1, 0]], dtype=torch.long) + with torch.no_grad(): + expected = reference(input_ids, attention_mask=attention_mask).last_hidden_state + actual = model(input_ids, attention_mask)["last_hidden_state"] + assert np.max(np.abs(expected.numpy() - actual.detach().numpy())) < 1e-5 + + +def test_graph_bert_backpropagates_through_attention_and_both_heads(): + GraphBERT = _load_model_module( + "graph_bert.py", "graph_bert_gradient_test").GraphBERT + model = GraphBERT( + vocab_size=24, + output_dim=2, + hidden_size=16, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=32, + max_position_embeddings=8, + dropout_rate=0.0, + task_dropout=0.0, + ) + input_ids = tlx.convert_to_tensor([[3, 5, 6, 4]], dtype=tlx.int64) + attention_mask = tlx.convert_to_tensor([[1, 1, 1, 1]], dtype=tlx.int64) + + loss = ( + model(input_ids, attention_mask=attention_mask, task="mlm").sum() + + model(input_ids, attention_mask=attention_mask, task="supervised").sum() + ) + loss.backward() + + checked = [ + model.encoder_layers[0].attention.q_proj.trainable_weights[0], + model.mlm_head.trainable_weights[0], + model.task_head_out.trainable_weights[0], + ] + for parameter in checked: + assert parameter.grad is not None + assert torch.isfinite(parameter.grad).all() + assert torch.count_nonzero(parameter.grad).item() > 0 + + +def test_graph_bert_attention_supports_nonsquare_sequence_and_head_shapes(): + GraphBERT = _load_model_module( + "graph_bert.py", "graph_bert_nonsquare_attention_test").GraphBERT + model = GraphBERT( + vocab_size=32, + output_dim=1, + hidden_size=24, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=48, + max_position_embeddings=16, + dropout_rate=0.0, + ) + input_ids = tlx.convert_to_tensor([[3, 5, 6, 7, 4]], dtype=tlx.int64) + + outputs = model(input_ids) + + assert tlx.get_tensor_shape(outputs["last_hidden_state"]) == [1, 5, 24] + + +def test_graph_gte_uses_long_position_window_and_forward_shape(): + graph_bert = _load_model_module("graph_bert.py", "graph_bert") + GraphGTE = _load_model_module("graph_gte.py", "graph_gte_under_test").GraphGTE + + model = GraphGTE( + vocab_size=48, + output_dim=1, + hidden_size=24, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=48, + max_position_embeddings=64, + ) + input_ids = tlx.convert_to_tensor([[3, 8, 9, 10, 4]], dtype=tlx.int64) + + outputs = model(input_ids) + + assert graph_bert.GraphBERT(vocab_size=48, output_dim=1).max_position_embeddings == 768 + assert model.max_position_embeddings == 64 + assert tlx.get_tensor_shape(outputs["logits"]) == [1, 1] + assert tlx.get_tensor_shape(outputs["mlm_logits"]) == [1, 5, 48] + + +def test_graph_gte_applies_rotary_embedding_to_query_and_key(): + _load_model_module("graph_bert.py", "graph_bert") + graph_gte = _load_model_module("graph_gte.py", "graph_gte_rope_test") + query = tlx.convert_to_tensor([[[[1.0, 2.0, 3.0, 4.0]]]]) + key = tlx.convert_to_tensor([[[[5.0, 6.0, 7.0, 8.0]]]]) + cosine = tlx.convert_to_tensor([[[[0.0, 0.0, 0.0, 0.0]]]]) + sine = tlx.convert_to_tensor([[[[1.0, 1.0, 1.0, 1.0]]]]) + + rotated_query, rotated_key = graph_gte.apply_rotary_pos_emb( + query, key, cosine, sine) + + assert torch.equal( + rotated_query, + torch.tensor([[[[-3.0, -4.0, 1.0, 2.0]]]]), + ) + assert torch.equal( + rotated_key, + torch.tensor([[[[-7.0, -8.0, 5.0, 6.0]]]]), + ) + + +def test_graph_gte_has_rope_without_absolute_position_embeddings(): + _load_model_module("graph_bert.py", "graph_bert") + GraphGTE = _load_model_module( + "graph_gte.py", "graph_gte_position_test").GraphGTE + model = GraphGTE( + vocab_size=32, + output_dim=2, + hidden_size=16, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=32, + max_position_embeddings=16, + ) + + assert not hasattr(model, "position_embeddings") + assert model.manifest()["position_embedding_type"] == "rope" + + +def test_graph_gte_manifest_reports_strict_paper_defaults(): + _load_model_module("graph_bert.py", "graph_bert") + GraphGTE = _load_model_module( + "graph_gte.py", "graph_gte_manifest_test").GraphGTE + model = GraphGTE(vocab_size=64, output_dim=1) + + manifest = model.manifest() + + assert manifest["encoder_type"] == "gte" + assert manifest["model_name"] == "gte-base" + assert manifest["weight_source"] == "random_initialization" + assert manifest["official_checkpoint"] is None + assert manifest["hidden_size"] == 768 + assert manifest["num_hidden_layers"] == 12 + assert manifest["num_attention_heads"] == 12 + assert manifest["intermediate_size"] == 3072 + assert manifest["hidden_act"] == "gelu" + assert manifest["hidden_dropout_prob"] == 0.1 + assert manifest["attention_probs_dropout_prob"] == 0.1 + assert manifest["position_embedding_type"] == "rope" + assert manifest["max_position_embeddings"] == 8192 + assert manifest["layer_norm_eps"] == 1e-12 + assert manifest["rope_theta"] == 20000.0 + assert manifest["rope_scaling"] == {"type": "ntk", "factor": 8.0} + assert 110_000_000 < manifest["total_parameters"] < 125_000_000 + assert manifest["trainable_parameters"] == manifest["total_parameters"] + + +def test_graph_gte_supports_paper_outputs_padding_and_gradients(): + _load_model_module("graph_bert.py", "graph_bert") + GraphGTE = _load_model_module( + "graph_gte.py", "graph_gte_training_test").GraphGTE + model = GraphGTE( + vocab_size=32, + output_dim=2, + hidden_size=16, + num_hidden_layers=1, + num_attention_heads=4, + intermediate_size=32, + max_position_embeddings=16, + dropout_rate=0.0, + ) + attention_mask = tlx.convert_to_tensor([[1, 1, 1, 0]], dtype=tlx.int64) + input_ids = tlx.convert_to_tensor([[3, 5, 4, 0]], dtype=tlx.int64) + + mlm_logits = model(input_ids, attention_mask=attention_mask, task="mlm") + task_logits = model( + input_ids, attention_mask=attention_mask, task="supervised") + + assert tlx.get_tensor_shape(mlm_logits) == [1, 4, 32] + assert tlx.get_tensor_shape(task_logits) == [1, 2] + loss = mlm_logits.sum() + task_logits.sum() + loss.backward() + checked = [ + model.encoder_layers[0].attention.qkv_proj.trainable_weights[0], + model.encoder_layers[0].mlp.up_gate_proj.trainable_weights[0], + model.mlm_head.trainable_weights[0], + model.task_head_out.trainable_weights[0], + ] + for parameter in checked: + assert parameter.grad is not None + assert torch.isfinite(parameter.grad).all() + assert torch.count_nonzero(parameter.grad).item() > 0 + + model.set_eval() + changed_padding = tlx.convert_to_tensor([[3, 5, 4, 17]], dtype=tlx.int64) + first = model(input_ids, attention_mask=attention_mask, task="supervised") + second = model( + changed_padding, attention_mask=attention_mask, task="supervised") + assert torch.allclose(first, second, atol=1e-6, rtol=1e-6) + + +def test_graph_tokenizer_trainer_resolves_paper_dataset_aliases(): + trainer = _load_trainer_module() + + qm9 = trainer.resolve_dataset_spec("QM9") + molhiv = trainer.resolve_dataset_spec("OGBG-molhiv") + peptides = trainer.resolve_dataset_spec("p-func") + peptides_struct = trainer.resolve_dataset_spec("Peptides-struct") + + assert qm9.canonical_name == "qm9" + assert qm9.task_type == "regression" + assert qm9.output_dim == 16 + assert qm9.metric == "mae" + assert molhiv.canonical_name == "molhiv" + assert molhiv.output_dim == 1 + assert molhiv.metric == "rocauc" + assert peptides.canonical_name == "peptides-func" + assert peptides.output_dim == 10 + assert peptides.metric == "ap" + assert peptides_struct.canonical_name == "peptides-struct" + assert peptides_struct.task_type == "multi_target_regression" + assert peptides_struct.output_dim == 11 + assert peptides_struct.metric == "average_mae" + + +def test_graph_tokenizer_trainer_pads_and_masks_token_batches(): + trainer = _load_trainer_module() + GraphTokenizer, _ = trainer.load_tokenizer_classes() + tokenizer = GraphTokenizer() + + batch = tokenizer.build_mlm_batch_from_token_sequences( + [[3, 6, 4], [3, 8, 9, 4]], max_length=5, mask_prob=1.0, seed=7) + + assert batch.input_ids == [[3, 2, 4, 0, 0], [3, 2, 2, 4, 0]] + assert batch.attention_mask == [[1, 1, 1, 0, 0], [1, 1, 1, 1, 0]] + assert batch.labels == [[-100, 6, -100, -100, -100], [-100, 8, 9, -100, -100]] + assert not hasattr(trainer, "pad_token_sequences") + assert not hasattr(trainer, "apply_mlm_mask") + + +def test_graph_tokenizer_trainer_builds_mlm_pretrain_batches(): + trainer = _load_trainer_module() + + encoded_split = { + "input_ids": [[3, 6, 4, 0], [3, 7, 8, 4]], + "attention_mask": [[1, 1, 1, 0], [1, 1, 1, 1]], + "labels": [[0.1], [0.2]], + } + GraphTokenizer, _ = trainer.load_tokenizer_classes() + batch = trainer.build_mlm_pretrain_split( + encoded_split, + tokenizer=GraphTokenizer(), + mask_prob=1.0, + seed=9, + ) + + assert batch == { + "input_ids": [[3, 2, 4, 0], [3, 7, 2, 4]], + "attention_mask": [[1, 1, 1, 0], [1, 1, 1, 1]], + "mlm_labels": [[-100, 6, -100, -100], [-100, -100, 8, -100]], + } + assert trainer.count_mlm_labels(batch["mlm_labels"]) == 2 + + +def test_graph_tokenizer_trainer_loads_preprocessed_split_files(): + trainer = _load_trainer_module() + spec = trainer.resolve_dataset_spec("qm9") + + graphs = [ + types.SimpleNamespace(edge_index=[[0], [1]], x=[index, index + 1], edge_attr=[index], y=[float(index)] * 16) + for index in range(3) + ] + class Dataset: + def __len__(self): + return len(graphs) + + def __getitem__(self, index): + return graphs[index] + + def get_idx_split(self): + return {"train": [0, 2], "val": [1], "test": []} + + dataset = Dataset() + trainer.load_gammagl_benchmark_dataset = lambda _root, _spec: dataset + + splits = trainer.load_benchmark_splits("unused", spec) + + assert [len(splits[name]) for name in ("train", "val", "test")] == [2, 1, 0] + assert splits["train"][0].edge_index == [[0], [1]] + assert splits["train"][0].x == [0, 1] + assert splits["train"][0].edge_attr == [0] + assert splits["train"][0].y[:2] == [0.0, 0.0] + assert [graph.graph_tokenizer_id for graph in splits["train"]] == [0, 2] + + +def test_graph_tokenizer_dataset_loader_bootstraps_repo_import_path(tmp_path, monkeypatch): + trainer = _load_trainer_module() + spec = trainer.resolve_dataset_spec("qm9") + calls = [] + fake_datasets = types.SimpleNamespace(QM9=lambda root: {"root": root}) + fake_gammagl = types.ModuleType("gammagl") + fake_gammagl.datasets = fake_datasets + + monkeypatch.setattr(trainer, "ensure_repo_on_path", lambda: calls.append(True)) + monkeypatch.setitem(sys.modules, "gammagl", fake_gammagl) + + dataset = trainer.load_gammagl_benchmark_dataset(tmp_path, spec) + + assert dataset == {"root": str(tmp_path)} + assert calls == [True] + + +def test_graph_tokenizer_bpe_preflight_bootstraps_repo_import_path(monkeypatch): + trainer = _load_trainer_module() + calls = [] + fake_third_party = types.ModuleType("third_party") + fake_third_party.graph_bpe_cpp = types.SimpleNamespace(is_available=lambda: True) + + monkeypatch.setattr(trainer, "ensure_repo_on_path", lambda: calls.append(True)) + monkeypatch.setitem(sys.modules, "third_party", fake_third_party) + + status = trainer.preflight_bpe_backend_status("cpp") + + assert status == {"backend": "cpp", "status": "ok", "native_available": True} + assert calls == [True] + + +def test_graph_tokenizer_trainer_rejects_benchmarks_without_a_public_dataset(): + trainer = _load_trainer_module() + with pytest.raises(ValueError, match="QM9"): + trainer.load_gammagl_benchmark_dataset("unused", trainer.resolve_dataset_spec("p-func")) + + +def test_graph_tokenizer_trainer_computes_primary_metrics(): + trainer = _load_trainer_module() + + qm9 = trainer.resolve_dataset_spec("qm9") + molhiv = trainer.resolve_dataset_spec("molhiv") + peptides = trainer.resolve_dataset_spec("p-func") + peptides_struct = trainer.resolve_dataset_spec("p-struct") + + assert trainer.compute_primary_metric(qm9, [[1.0, 2.0]], [[2.0, 4.0]]) == 1.5 + assert trainer.compute_primary_metric(molhiv, [0, 1, 0, 1], [0.1, 0.9, 0.4, 0.8]) == 1.0 + assert trainer.compute_primary_metric(peptides, [[1, 0], [0, 1]], [[0.9, 0.1], [0.2, 0.8]]) == 1.0 + assert trainer.compute_primary_metric( + peptides_struct, + [[1.0, 2.0], [3.0, float("nan")]], + [[2.0, 4.0], [1.0, 100.0]], + ) == 1.75 + + +def test_shared_mlm_helper_forces_one_mask_when_random_draw_selects_none(): + GraphTokenizer, _ = _load_trainer_module().load_tokenizer_classes() + masked, labels = GraphTokenizer().mask_token_sequences( + [[3, 11, 12, 4]], mask_prob=1e-12, seed=0) + + selected = [ + index for index, label in enumerate(labels[0]) if label != -100 + ] + assert len(selected) == 1 + assert labels[0][selected[0]] in {11, 12} + assert masked[0][selected[0]] == 2 + + +def test_graph_tokenizer_trainer_reports_task_metrics_by_dataset(): + trainer = _load_trainer_module() + + qm9 = trainer.resolve_dataset_spec("qm9") + molhiv = trainer.resolve_dataset_spec("molhiv") + peptides_func = trainer.resolve_dataset_spec("p-func") + peptides_struct = trainer.resolve_dataset_spec("p-struct") + + qm9_metrics = trainer.compute_task_metrics(qm9, [[1.0, 2.0]], [[2.0, 4.0]]) + molhiv_metrics = trainer.compute_task_metrics(molhiv, [0, 1, 0, 1], [0.1, 0.9, 0.4, 0.8]) + peptides_func_metrics = trainer.compute_task_metrics( + peptides_func, + [[1, 0], [0, 1]], + [[0.9, 0.1], [0.2, 0.8]], + ) + peptides_struct_metrics = trainer.compute_task_metrics( + peptides_struct, + [[1.0, 2.0], [3.0, float("nan")]], + [[2.0, 4.0], [1.0, 100.0]], + ) + + assert qm9_metrics == {"mae": 1.5} + assert molhiv_metrics == {"auc": 1.0} + assert peptides_func_metrics == {"ap": 1.0} + assert peptides_struct_metrics == {"average_mae": 1.75, "mae": 1.6666666666666667} + assert trainer.accuracy_score([0, 1, 1, 0], [0.2, 0.8, 0.7, 0.4]) == 1.0 + + +def test_graph_tokenizer_trainer_selects_best_epoch_by_metric_direction(): + trainer = _load_trainer_module() + + qm9 = trainer.resolve_dataset_spec("qm9") + molhiv = trainer.resolve_dataset_spec("molhiv") + + regression_history = [ + {"epoch": 1, "val": {"metrics": {"mae": 0.9}}, "test": {"metrics": {"mae": 1.1}}}, + {"epoch": 2, "val": {"metrics": {"mae": 0.7}}, "test": {"metrics": {"mae": 1.0}}}, + {"epoch": 3, "val": {"metrics": {"mae": 0.8}}, "test": {"metrics": {"mae": 0.6}}}, + ] + classification_history = [ + {"epoch": 1, "val": {"metrics": {"auc": 0.71}}, "test": {"metrics": {"auc": 0.69}}}, + {"epoch": 2, "val": {"metrics": {"auc": 0.68}}, "test": {"metrics": {"auc": 0.80}}}, + {"epoch": 3, "val": {"metrics": {"auc": 0.74}}, "test": {"metrics": {"auc": 0.73}}}, + ] + + assert trainer.select_best_epoch(qm9, regression_history)["epoch"] == 2 + assert trainer.select_best_epoch(molhiv, classification_history)["epoch"] == 3 + + +def test_test_evaluated_after_checkpoint_selection(monkeypatch, tmp_path): + trainer = _load_trainer_module() + spec = trainer.resolve_dataset_spec("qm9") + splits = {"train": [object()], "val": [object()], "test": [object()]} + encoded = {name: {"split": name} for name in splits} + calls = [] + + class Model: + trainable_weights = [] + + def __init__(self, **_kwargs): + self.loaded_state = None + + def state_dict(self): + return {"epoch": len(calls)} + + def load_state_dict(self, state): + self.loaded_state = state + + model = Model() + fake_tlx = types.ModuleType("tensorlayerx") + fake_tlx.optimizers = SimpleNamespace(Adam=lambda **_kwargs: object()) + fake_tlx_model = types.ModuleType("tensorlayerx.model") + fake_tlx_model.TrainOneStep = lambda *_args: object() + monkeypatch.setitem(sys.modules, "tensorlayerx", fake_tlx) + monkeypatch.setitem(sys.modules, "tensorlayerx.model", fake_tlx_model) + monkeypatch.setattr(trainer, "load_benchmark_splits", lambda *_args: splits) + monkeypatch.setattr(trainer, "fit_tokenizer", lambda *_args: object()) + monkeypatch.setattr(trainer, "encode_graph_splits", lambda *_args, **_kwargs: encoded) + monkeypatch.setattr(trainer, "load_model_class", lambda _name: lambda **_kwargs: model) + monkeypatch.setattr(trainer, "model_kwargs", lambda *_args: {}) + monkeypatch.setattr(trainer, "run_mlm_pretrain", lambda *_args: {}) + monkeypatch.setattr(trainer, "GraphTokenSupervisedLoss", lambda *_args: object()) + trained_splits = [] + monkeypatch.setattr( + trainer, "train_one_epoch", + lambda _model, _step, encoded_train, *_args: ( + trained_splits.append(encoded_train["split"]) or 0.0), + ) + + val_metrics = iter((0.4, 0.2, 0.3)) + + def evaluate(_model, encoded_split, _spec, _args): + split = encoded_split["split"] + calls.append(split) + metric = next(val_metrics) if split == "val" else 0.1 + return {"loss": metric, "metrics": {"mae": metric}} + + monkeypatch.setattr(trainer, "evaluate_split", evaluate) + args = SimpleNamespace( + dataset="qm9", seed=0, data_root=tmp_path, max_length=8, model="bert", + pretrain_epoch=0, lr=0.001, weight_decay=0.0, n_epoch=3, patience=0, + batch_size=1, save_checkpoint=False, run=0, output_dir=tmp_path, + ) + + summary = trainer.run_train_val_test(args) + + assert calls == ["val", "val", "val", "test"] + assert trained_splits == ["train", "train", "train"] + assert model.loaded_state == {"epoch": 2} + assert all("test" not in entry for entry in summary["history"]) + assert summary["best_test"]["metrics"] == {"mae": 0.1} + + +def test_bpe_fit_uses_train_only(tmp_path): + trainer = _load_trainer_module() + args = types.SimpleNamespace( + protocol="gammagl", dataset="qm9", data_root=str(tmp_path), + serialization="feuler", num_merges=1, min_frequency=1, + bpe_backend="python") + train_graph = trainer.SyntheticGraph( + edge_index=[[0], [1]], x=[1, 2], edge_attr=[3], y=[0.0]) + held_out_graph = trainer.SyntheticGraph( + edge_index=[[0], [1]], x=[101, 102], edge_attr=[103], y=[1.0]) + tokenizer = trainer.fit_tokenizer(args, [train_graph]) + + held_out_pattern = (101, 103, 102) + assert held_out_pattern not in tokenizer.serializer.frequency_map + assert all(held_out_pattern[0] not in rule + for rule in tokenizer.bpe.codebook.merge_rules) + + +def test_graph_tokenizer_trainer_iterates_token_batches_with_labels(): + trainer = _load_trainer_module() + + batches = list( + trainer.iter_token_batches( + { + "input_ids": [[1, 2], [3, 4], [5, 6]], + "attention_mask": [[1, 1], [1, 1], [1, 0]], + "labels": [[0.1], [0.2], [0.3]], + }, + batch_size=2, + ) + ) + + assert batches == [ + { + "input_ids": [[1, 2], [3, 4]], + "attention_mask": [[1, 1], [1, 1]], + "labels": [[0.1], [0.2]], + }, + { + "input_ids": [[5, 6]], + "attention_mask": [[1, 0]], + "labels": [[0.3]], + }, + ] + + +def test_graph_tokenizer_trainer_rejects_multidimensional_features_without_adapter(): + trainer = _load_trainer_module() + + with pytest.raises(ValueError, match="dataset-specific token adapter"): + trainer.flatten_feature_ids([[1, 2], [1, 9]]) + + +def test_graph_tokenizer_trainer_formats_result_rows_and_aggregates_runs(): + trainer = _load_trainer_module() + spec = trainer.resolve_dataset_spec("molhiv") + summaries = [ + { + "dataset": "molhiv", + "model": "bert", + "run": 0, + "seed": 7, + "primary_metric": "auc", + "best_epoch": 2, + "best_val": {"loss": 0.4, "metrics": {"auc": 0.8}}, + "best_test": {"loss": 0.5, "metrics": {"auc": 0.7}}, + }, + { + "dataset": "molhiv", + "model": "bert", + "run": 1, + "seed": 8, + "primary_metric": "auc", + "best_epoch": 3, + "best_val": {"loss": 0.3, "metrics": {"auc": 0.9}}, + "best_test": {"loss": 0.6, "metrics": {"auc": 0.9}}, + }, + ] + + assert trainer.format_result_row(summaries[0]) == { + "dataset": "molhiv", + "model": "bert", + "run": 0, + "seed": 7, + "best_epoch": 2, + "val_auc": 0.8, + "test_auc": 0.7, + "val_loss": 0.4, + "test_loss": 0.5, + } + assert trainer.aggregate_run_summaries(spec, summaries) == { + "num_runs": 2, + "primary_metric": "auc", + "higher_is_better": True, + "val_auc_mean": 0.8500000000000001, + "val_auc_std": 0.04999999999999999, + "test_auc_mean": 0.8, + "test_auc_std": 0.10000000000000003, + } + + +def test_graph_tokenizer_aggregates_per_target_mae_across_runs(): + trainer = _load_trainer_module() + runs = [ + { + "best_val": {"per_target_mae": {"mu": 1.0, "alpha": 3.0}}, + "best_test": {"per_target_mae": {"mu": 2.0, "alpha": 4.0}}, + }, + { + "best_val": {"per_target_mae": {"mu": 3.0, "alpha": 5.0}}, + "best_test": {"per_target_mae": {"mu": 4.0, "alpha": 6.0}}, + }, + ] + + aggregate = trainer.aggregate_per_target_mae(runs) + + assert aggregate == { + "val": { + "mu": {"mean": 2.0, "std": 1.0}, + "alpha": {"mean": 4.0, "std": 1.0}, + }, + "test": { + "mu": {"mean": 3.0, "std": 1.0}, + "alpha": {"mean": 5.0, "std": 1.0}, + }, + } + + +def test_graph_tokenizer_aggregates_raw_per_target_mae_across_runs(): + trainer = _load_trainer_module() + runs = [ + { + "best_val": {"per_target_mae_raw": {"mu": 10.0}}, + "best_test": {"per_target_mae_raw": {"mu": 20.0}}, + }, + { + "best_val": {"per_target_mae_raw": {"mu": 30.0}}, + "best_test": {"per_target_mae_raw": {"mu": 40.0}}, + }, + ] + + assert trainer.aggregate_per_target_mae( + runs, field_name="per_target_mae_raw") == { + "val": {"mu": {"mean": 20.0, "std": 10.0}}, + "test": {"mu": {"mean": 30.0, "std": 10.0}}, + } + + +def test_graph_tokenizer_trainer_formats_paper_result_tables(): + trainer = _load_trainer_module() + + aggregate = { + "num_runs": 3, + "primary_metric": "auc", + "higher_is_better": True, + "val_auc_mean": 0.81234, + "val_auc_std": 0.01234, + "test_auc_mean": 0.76543, + "test_auc_std": 0.02345, + } + summary = { + "dataset": "molhiv", + "model": "bert", + "aggregate": aggregate, + } + + table_row = trainer.format_paper_table_row(summary) + + assert table_row == { + "dataset": "molhiv", + "model": "bert", + "metric": "auc", + "direction": "higher", + "runs": 3, + "val": "0.8123 +/- 0.0123", + "test": "0.7654 +/- 0.0234", + } + assert trainer.format_markdown_result_table([summary]).splitlines() == [ + "| Dataset | Model | Metric | Direction | Runs | Val | Test |", + "|---|---|---|---|---:|---:|---:|", + "| molhiv | bert | auc | higher | 3 | 0.8123 +/- 0.0123 | 0.7654 +/- 0.0234 |", + ] + assert trainer.format_latex_result_rows([summary]) == ( + "molhiv & bert & auc & higher & 3 & 0.8123 $\\pm$ 0.0123 & 0.7654 $\\pm$ 0.0234 \\" + ) + + +def test_graph_tokenizer_trainer_parses_experiment_matrix(): + trainer = _load_trainer_module() + + assert trainer.parse_csv_values("qm9, molhiv,,p-func") == ["qm9", "molhiv", "p-func"] + + class Args: + dataset = "qm9" + datasets = "qm9,p-struct" + model = "bert" + models = "bert,gte" + + assert trainer.experiment_matrix_items(Args()) == [ + ("qm9", "bert"), + ("qm9", "gte"), + ("p-struct", "bert"), + ("p-struct", "gte"), + ] + + +def test_graph_tokenizer_trainer_formats_combined_paper_tables(): + trainer = _load_trainer_module() + + summaries = [ + { + "dataset": "qm9", + "model": "bert", + "aggregate": { + "num_runs": 2, + "primary_metric": "mae", + "higher_is_better": False, + "val_mae_mean": 0.21, + "val_mae_std": 0.01, + "test_mae_mean": 0.23, + "test_mae_std": 0.02, + }, + }, + { + "dataset": "molhiv", + "model": "gte", + "aggregate": { + "num_runs": 2, + "primary_metric": "auc", + "higher_is_better": True, + "val_auc_mean": 0.81, + "val_auc_std": 0.03, + "test_auc_mean": 0.78, + "test_auc_std": 0.04, + }, + }, + ] + + markdown = trainer.format_markdown_result_table(summaries) + latex = trainer.format_latex_result_rows(summaries) + + assert "| qm9 | bert | mae | lower | 2 | 0.2100 +/- 0.0100 | 0.2300 +/- 0.0200 |" in markdown + assert "| molhiv | gte | auc | higher | 2 | 0.8100 +/- 0.0300 | 0.7800 +/- 0.0400 |" in markdown + assert "qm9 & bert & mae & lower & 2 & 0.2100 $\\pm$ 0.0100 & 0.2300 $\\pm$ 0.0200" in latex + assert "molhiv & gte & auc & higher & 2 & 0.8100 $\\pm$ 0.0300 & 0.7800 $\\pm$ 0.0400" in latex + + +def test_graph_tokenizer_trainer_loads_existing_experiment_for_resume(tmp_path): + trainer = _load_trainer_module() + + class Args: + output_dir = str(tmp_path) + + summary_path = tmp_path / "qm9_bert_summary.json" + summary_path.write_text( + json.dumps({"dataset": "qm9", "model": "bert", "aggregate": {"num_runs": 1}}), + encoding="utf-8", + ) + + loaded = trainer.load_existing_experiment_summary(Args(), "qm9", "bert") + + assert loaded["dataset"] == "qm9" + assert loaded["model"] == "bert" + assert loaded["status"] == "skipped_existing" + assert loaded["outputs"]["json"] == str(summary_path) + + +def test_graph_tokenizer_trainer_records_matrix_failures(): + trainer = _load_trainer_module() + + failure = trainer.format_failed_experiment_summary("qm9", "bert", RuntimeError("boom")) + + assert failure["dataset"] == "qm9" + assert failure["model"] == "bert" + assert failure["status"] == "failed" + assert failure["error_type"] == "RuntimeError" + assert failure["error"] == "boom" + assert failure["aggregate"]["num_runs"] == 0 + + +def test_graph_tokenizer_trainer_serializes_args_for_manifest(tmp_path): + trainer = _load_trainer_module() + + class Args: + dataset = "qm9" + model = "bert" + output_dir = tmp_path + hidden_size = None + runs = 3 + values = ("a", "b") + + serialized = trainer.serialize_args(Args()) + + assert serialized["dataset"] == "qm9" + assert serialized["model"] == "bert" + assert serialized["output_dir"] == str(tmp_path) + assert serialized["hidden_size"] is None + assert serialized["runs"] == 3 + assert serialized["values"] == ["a", "b"] + + +def test_graph_tokenizer_trainer_builds_runtime_manifest(): + trainer = _load_trainer_module() + + class Args: + dataset = "molhiv" + model = "gte" + runs = 2 + output_dir = "logs" + + manifest = trainer.build_runtime_manifest(Args(), timestamp="2026-07-03T00:00:00+00:00") + + assert manifest["timestamp"] == "2026-07-03T00:00:00+00:00" + assert manifest["args"]["dataset"] == "molhiv" + assert manifest["args"]["model"] == "gte" + assert "version" in manifest["python"] + assert "executable" in manifest["python"] + assert "platform" in manifest["python"] + assert "packages" in manifest + assert "git" in manifest + + +def test_runtime_manifest_does_not_execute_git_commands(): + trainer = _load_trainer_module() + + assert trainer.git_metadata() == { + "status": "disabled", + "reason": "local-only run; git commands are disabled", + } + + +def test_graph_tokenizer_trainer_writes_manifest_output(tmp_path): + trainer = _load_trainer_module() + + class Args: + output_dir = str(tmp_path) + + summary = { + "dataset": "qm9", + "model": "bert", + "rows": [], + "aggregate": { + "num_runs": 1, + "primary_metric": "mae", + "higher_is_better": False, + "val_mae_mean": 0.2, + "val_mae_std": 0.0, + "test_mae_mean": 0.3, + "test_mae_std": 0.0, + }, + "manifest": {"timestamp": "2026-07-03T00:00:00+00:00", "args": {"dataset": "qm9"}}, + } + + outputs = trainer.write_experiment_outputs(summary, Args()) + manifest_path = Path(outputs["manifest"]) + + assert manifest_path.name == "qm9_bert_manifest.json" + assert json.loads(manifest_path.read_text(encoding="utf-8"))["timestamp"] == "2026-07-03T00:00:00+00:00" + + + +def test_graph_tokenizer_trainer_parses_paper_protocol_options(): + trainer = _load_trainer_module() + + args = trainer.parse_args([ + "--protocol", "paper", + "--dataset", "qm9", + "--model", "gte", + "--target-property", "homo", + "--device", "cuda:0", + "--pretrain-epoch", "200", + "--n-epoch", "200", + "--pretrain-lr", "0.00005", + "--lr", "0.00001", + "--batch-size", "32", + "--weight-decay", "0.1", + "--pretrain-warmup-ratio", "0.12", + "--finetune-warmup-ratio", "0.025", + "--pretrain-max-grad-norm", "2.0", + "--finetune-max-grad-norm", "0.5", + "--mask-prob", "0.09", + "--patience", "20", + "--max-position-embeddings", "8192", + "--paper-amp", "bf16", + "--paper-tf32", + "--paper-loss", "huber", + "--molhiv-pos-weight", + ]) + + assert args.protocol == "paper" + assert args.target_property == "homo" + assert args.device == "cuda:0" + assert args.pretrain_epoch == 200 + assert args.n_epoch == 200 + assert args.pretrain_lr == 5e-5 + assert args.lr == 1e-5 + assert args.batch_size == 32 + assert args.weight_decay == 0.1 + assert args.pretrain_warmup_ratio == 0.12 + assert args.finetune_warmup_ratio == 0.025 + assert args.pretrain_max_grad_norm == 2.0 + assert args.finetune_max_grad_norm == 0.5 + assert args.mask_prob == 0.09 + assert args.patience == 20 + assert args.max_position_embeddings == 8192 + assert args.paper_amp == "bf16" + assert args.paper_tf32 is True + assert args.paper_loss == "huber" + assert args.molhiv_pos_weight is True + + + +def test_graph_tokenizer_trainer_rejects_unknown_model_name(): + trainer = _load_trainer_module() + + try: + trainer.resolve_model_name("bad-model") + except ValueError as error: + assert "Unsupported model" in str(error) + else: + raise AssertionError("unknown model name was accepted") + + +def test_graph_tokenizer_trainer_rejects_overlapping_dataset_splits(monkeypatch): + trainer = _load_trainer_module() + class Dataset(list): + def get_idx_split(self): + return {"train": [0], "val": [0], "test": [1]} + + dataset = Dataset([ + types.SimpleNamespace(edge_index=[[], []], x=[], edge_attr=[], y=[0.0] * 16) + for _ in range(2) + ]) + monkeypatch.setattr(trainer, "load_gammagl_benchmark_dataset", lambda *_args: dataset) + + with pytest.raises(ValueError, match="overlaps"): + trainer.load_benchmark_splits("unused", trainer.resolve_dataset_spec("qm9")) + + +def test_paper_tokenizer_reuses_train_only_local_cache(tmp_path): + trainer = _load_trainer_module() + args = types.SimpleNamespace( + protocol="paper", + dataset="qm9", + data_root=str(tmp_path / "data"), + paper_cache_root=str(tmp_path / "cache"), + serialization="feuler", + num_merges=1, + min_frequency=1, + bpe_backend="python", + ) + graphs = [ + trainer.SyntheticGraph( + edge_index=[[0], [1]], x=[13, 15], edge_attr=[2], y=[0.0]), + trainer.SyntheticGraph( + edge_index=[[0], [1]], x=[13, 17], edge_attr=[2], y=[1.0]), + ] + + first = trainer.fit_tokenizer(args, graphs) + second = trainer.fit_tokenizer(args, graphs) + + assert first._cache_status == "miss" + assert second._cache_status == "hit" + assert second.bpe.codebook.merge_rules == first.bpe.codebook.merge_rules + assert len(list((tmp_path / "cache").rglob("tokenizer.pkl"))) == 1 + + +def test_graph_tokenizer_trainer_preflight_reports_missing_dataset(tmp_path, monkeypatch): + trainer = _load_trainer_module() + monkeypatch.setattr( + trainer, "load_gammagl_benchmark_dataset", lambda data_root, spec: None) + + class Args: + data_root = str(tmp_path) + dataset = "qm9" + datasets = "qm9" + model = "bert" + models = "bert" + bpe_backend = "python" + + report = trainer.preflight_check(Args()) + + assert report["status"] == "failed" + assert report["num_errors"] == 1 + assert report["datasets"][0]["status"] == "failed" + assert "dataset loader could not materialize" in report["errors"][0] diff --git a/tests/transforms/test_graph_bpe_cpp_install.py b/tests/transforms/test_graph_bpe_cpp_install.py new file mode 100644 index 000000000..909a395d2 --- /dev/null +++ b/tests/transforms/test_graph_bpe_cpp_install.py @@ -0,0 +1,38 @@ +import ast +import importlib +from pathlib import Path + +import pytest + + +def test_graph_bpe_setup_targets_runtime_module_name(): + setup_path = ( + Path(__file__).resolve().parents[2] + / "third_party" + / "graph_bpe_cpp" + / "setup.py" + ) + tree = ast.parse(setup_path.read_text(encoding="utf-8")) + extension_names = [ + node.args[0].value + for node in ast.walk(tree) + if isinstance(node, ast.Call) + and isinstance(node.func, ast.Name) + and node.func.id == "Extension" + and node.args + and isinstance(node.args[0], ast.Constant) + ] + + assert extension_names == ["third_party.graph_bpe_cpp._graph_bpe"] + + +def test_built_graph_bpe_extension_matches_python_reference(): + native = pytest.importorskip("third_party.graph_bpe_cpp._graph_bpe") + bridge = importlib.import_module("third_party.graph_bpe_cpp") + sequences = [[1, 2, 1, 2], [1, 2, 1, 2], [1, 2, 3]] + + native_result = native.train_bpe(sequences, 2, 2) + reference = bridge._train_bpe_python(sequences, 2, 2) + + assert native_result["merge_rules"] == reference["merge_rules"] + assert native_result["vocab_size"] == reference["vocab_size"] diff --git a/tests/transforms/test_graph_tokenizer.py b/tests/transforms/test_graph_tokenizer.py new file mode 100644 index 000000000..93e37f036 --- /dev/null +++ b/tests/transforms/test_graph_tokenizer.py @@ -0,0 +1,1007 @@ +import importlib.util +import random +import sys +import types +from pathlib import Path + +import pytest + +def _load_graph_serializer_module(): + module_path = Path(__file__).resolve().parents[2] / "gammagl" / "transforms" / "graph_serializer.py" + spec = importlib.util.spec_from_file_location("graph_serializer_under_test", module_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def _load_graph_bpe_module(): + module_path = Path(__file__).resolve().parents[2] / "gammagl" / "transforms" / "graph_bpe.py" + spec = importlib.util.spec_from_file_location("graph_bpe_under_test", module_path) + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module + + +def _load_graph_tokenizer_module(): + root = Path(__file__).resolve().parents[2] + gammagl_package = sys.modules.setdefault("gammagl", types.ModuleType("gammagl")) + gammagl_package.__path__ = [str(root / "gammagl")] + transforms_package = sys.modules.setdefault("gammagl.transforms", types.ModuleType("gammagl.transforms")) + transforms_package.__path__ = [str(root / "gammagl" / "transforms")] + module_path = Path(__file__).resolve().parents[2] / "gammagl" / "transforms" / "graph_tokenizer.py" + spec = importlib.util.spec_from_file_location("gammagl.transforms.graph_tokenizer_under_test", module_path) + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +def test_graph_tokenizer_public_interfaces_are_importable(): + graph_bpe = _load_graph_bpe_module() + graph_serializer = _load_graph_serializer_module() + graph_tokenizer = _load_graph_tokenizer_module() + BPECodebook = graph_bpe.BPECodebook + GraphBPE = graph_bpe.GraphBPE + FrequencyGuidedEulerianSerializer = graph_serializer.FrequencyGuidedEulerianSerializer + GraphTokenizer = graph_tokenizer.GraphTokenizer + GraphTokenizationResult = graph_tokenizer.GraphTokenizationResult + GraphTokenizerMLMBatch = graph_tokenizer.GraphTokenizerMLMBatch + from third_party.graph_bpe_cpp import is_available + + serializer = FrequencyGuidedEulerianSerializer() + bpe = GraphBPE() + tokenizer = GraphTokenizer(serializer=serializer, bpe=bpe) + codebook = BPECodebook(merge_rules=[], vocab_size=0) + + assert serializer.name == "feuler" + assert bpe.backend == "python" + assert tokenizer.serializer is serializer + assert tokenizer.bpe is bpe + assert codebook.merge_rules == [] + assert GraphTokenizationResult(input_ids=[], attention_mask=[], serialized_token_ids=[], metadata={}).input_ids == [] + assert GraphTokenizerMLMBatch(input_ids=[], attention_mask=[], labels=[], metadata={}).labels == [] + assert isinstance(is_available(), bool) + + +class SimpleGraph: + def __init__(self, edge_index, x, edge_attr=None, num_nodes=None): + self.edge_index = edge_index + self.x = x + self.edge_attr = edge_attr + self.num_nodes = num_nodes if num_nodes is not None else len(x) + + +def test_feuler_fit_counts_node_edge_node_patterns(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + + graph = SimpleGraph( + edge_index=[[0, 1, 1], [1, 0, 2]], + x=[7, 8, 9], + edge_attr=[3, 3, 4], + ) + + serializer = FrequencyGuidedEulerianSerializer() + serializer.fit([graph]) + + assert serializer.frequency_map == { + (7, 3, 8): 1, + (8, 4, 9): 1, + } + + +def test_single_vs_symmetric_coo_same_frequency(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + single_direction = SimpleGraph( + edge_index=[[0], [1]], x=[7, 8], edge_attr=[3]) + symmetric = SimpleGraph( + edge_index=[[0, 1], [1, 0]], x=[7, 8], edge_attr=[3, 3]) + + single_serializer = FrequencyGuidedEulerianSerializer().fit([single_direction]) + symmetric_serializer = FrequencyGuidedEulerianSerializer().fit([symmetric]) + + assert single_serializer.frequency_map == symmetric_serializer.frequency_map + + +def test_single_vs_symmetric_coo_same_serialization(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + single_direction = SimpleGraph( + edge_index=[[0], [1]], x=[7, 8], edge_attr=[3]) + symmetric = SimpleGraph( + edge_index=[[0, 1], [1, 0]], x=[7, 8], edge_attr=[3, 3]) + + single_serializer = FrequencyGuidedEulerianSerializer().fit([single_direction]) + symmetric_serializer = FrequencyGuidedEulerianSerializer().fit([symmetric]) + + assert single_serializer.serialize(single_direction).token_ids == ( + symmetric_serializer.serialize(symmetric).token_ids) + + +def test_invalid_node_feature_length_raises(): + serializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer() + graph = SimpleGraph( + edge_index=[[0], [1]], x=[7], edge_attr=[3], num_nodes=2) + + with pytest.raises(ValueError, match="x must contain exactly"): + serializer.fit([graph]) + + +def test_invalid_edge_feature_length_raises(): + serializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer() + graph = SimpleGraph( + edge_index=[[0, 1], [1, 2]], x=[7, 8, 9], edge_attr=[3]) + + with pytest.raises(ValueError, match="edge_attr must contain exactly"): + serializer.fit([graph]) + + +def test_missing_node_features_supported(): + serializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer() + graph = SimpleGraph( + edge_index=[[0], [1]], x=None, edge_attr=[3], num_nodes=2) + + result = serializer.fit([graph]).serialize(graph) + + assert result.token_ids == [ + serializer.node_token(0), serializer.node_reference_token(0), + serializer.edge_token(3), + serializer.node_token(1), serializer.node_reference_token(1), + serializer.edge_token(3), + serializer.node_token(0), serializer.node_reference_token(0), + ] + + +def test_missing_edge_features_supported(): + serializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer() + graph = SimpleGraph( + edge_index=[[0], [1]], x=[7, 8], edge_attr=None) + + result = serializer.fit([graph]).serialize(graph) + + assert result.token_ids == [ + serializer.node_token(7), serializer.node_reference_token(0), + serializer.edge_token(0), + serializer.node_token(8), serializer.node_reference_token(1), + serializer.edge_token(0), + serializer.node_token(7), serializer.node_reference_token(0), + ] + + +def test_feuler_serialization_prefers_high_frequency_edge(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + + training_graphs = [ + SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]), + SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]), + SimpleGraph(edge_index=[[0], [1]], x=[1, 3], edge_attr=[9]), + ] + target = SimpleGraph(edge_index=[[0, 0], [1, 2]], x=[1, 2, 3], edge_attr=[5, 9]) + + serializer = FrequencyGuidedEulerianSerializer() + serializer.fit(training_graphs) + result = serializer.serialize(target) + + node_one = serializer.node_token(1) + node_two = serializer.node_token(2) + edge_five = serializer.edge_token(5) + + assert result.token_ids[:5] == [ + node_one, + serializer.node_reference_token(0), + edge_five, + node_two, + serializer.node_reference_token(1), + ] + assert result.metadata["method"] == "feuler" + assert result.metadata["num_edges_traversed"] == 4 + + +def test_feuler_serialization_sorts_and_joins_components(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + + graph = SimpleGraph(edge_index=[[0, 2], [1, 3]], x=[4, 5, 6, 7], edge_attr=[1, 2]) + serializer = FrequencyGuidedEulerianSerializer(component_sep_token_id=-1) + + result = serializer.fit([graph]).serialize(graph) + + assert -1 in result.token_ids + assert result.metadata["num_components"] == 2 + + +def test_feuler_serializes_long_graph_without_python_recursion(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + num_edges = 1200 + graph = SimpleGraph( + edge_index=[list(range(num_edges)), list(range(1, num_edges + 1))], + x=[index % 5 for index in range(num_edges + 1)], + edge_attr=[1] * num_edges, + ) + + result = FrequencyGuidedEulerianSerializer().fit([graph]).serialize(graph) + + assert result.metadata["num_edges_traversed"] == num_edges * 2 + assert len(result.token_ids) == num_edges * 6 + 2 + + +def test_feuler_deserializes_from_token_stream_without_graph_metadata(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + + graph = SimpleGraph(edge_index=[[0, 1], [1, 2]], x=[4, 5, 6], edge_attr=[1, 2]) + serializer = FrequencyGuidedEulerianSerializer().fit([graph]) + + result = serializer.serialize(graph) + restored = serializer.deserialize(result) + + assert restored["edge_index"] == [[0, 1], [1, 2]] + assert restored["x"] == [4, 5, 6] + assert restored["edge_attr"] == [1, 2] + assert "node_labels" not in result.metadata + assert "input_edges" not in result.metadata + + +def _canonical_labeled_edges(edge_index, edge_attr): + labels = edge_attr if edge_attr is not None else [0] * len(edge_index[0]) + return sorted( + (min(src, dst), max(src, dst), label) + for src, dst, label in zip(edge_index[0], edge_index[1], labels)) + + +def _component_sets(num_nodes, edge_index): + neighbors = {node: set() for node in range(num_nodes)} + for src, dst in zip(edge_index[0], edge_index[1]): + neighbors[src].add(dst) + neighbors[dst].add(src) + components = [] + remaining = set(range(num_nodes)) + while remaining: + pending = [min(remaining)] + component = set() + while pending: + node = pending.pop() + if node in component: + continue + component.add(node) + pending.extend(neighbors[node] - component) + remaining -= component + components.append(frozenset(component)) + return set(components) + + +def _assert_feuler_round_trip(graph): + serializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer().fit([graph]) + reconstructed = serializer.deserialize(serializer.serialize(graph)) + input_x = graph.x if graph.x is not None else list(range(graph.num_nodes)) + + assert reconstructed["num_nodes"] == graph.num_nodes + assert reconstructed["x"] == input_x + assert _canonical_labeled_edges( + reconstructed["edge_index"], reconstructed["edge_attr"]) == _canonical_labeled_edges( + graph.edge_index, graph.edge_attr) + assert _component_sets( + reconstructed["num_nodes"], reconstructed["edge_index"]) == _component_sets( + graph.num_nodes, graph.edge_index) + + +def test_round_trip_single_edge(): + _assert_feuler_round_trip(SimpleGraph( + edge_index=[[0], [1]], x=[4, 9], edge_attr=[7])) + + +def test_round_trip_path(): + _assert_feuler_round_trip(SimpleGraph( + edge_index=[[0, 1], [1, 2]], x=[4, 9, 4], edge_attr=[7, 8])) + + +def test_round_trip_cycle(): + _assert_feuler_round_trip(SimpleGraph( + edge_index=[[0, 1, 2], [1, 2, 0]], x=[1, 1, 1], edge_attr=[2, 3, 4])) + + +def test_round_trip_star(): + _assert_feuler_round_trip(SimpleGraph( + edge_index=[[0, 0, 0], [1, 2, 3]], x=[5, 5, 5, 5], edge_attr=[1, 2, 3])) + + +def test_round_trip_disconnected_graph(): + _assert_feuler_round_trip(SimpleGraph( + edge_index=[[0, 2], [1, 3]], x=[1, 2, 3, 4], edge_attr=[5, 6])) + + +def test_round_trip_isolated_nodes(): + _assert_feuler_round_trip(SimpleGraph( + edge_index=[[0], [1]], x=[1, 2, 3, 4], edge_attr=[5], num_nodes=4)) + + +def test_round_trip_node_labels(): + _assert_feuler_round_trip(SimpleGraph( + edge_index=[[0, 1], [1, 2]], x=[-3, 0, 11], edge_attr=[5, 5])) + + +def test_round_trip_edge_labels(): + _assert_feuler_round_trip(SimpleGraph( + edge_index=[[0, 1], [1, 2]], x=[1, 2, 3], edge_attr=[-4, 9])) + + +def _feuler_structure_tokens(token_ids): + """Remove reversible node references before comparing traversals.""" + return [token for token in token_ids if token == -1 or token % 3 != 2] + + +def _relabel_graph(graph, old_to_new): + num_nodes = graph.num_nodes + new_to_old = {new: old for old, new in old_to_new.items()} + return SimpleGraph( + edge_index=[ + [old_to_new[node] for node in graph.edge_index[0]], + [old_to_new[node] for node in graph.edge_index[1]], + ], + x=[graph.x[new_to_old[node]] for node in range(num_nodes)], + edge_attr=list(graph.edge_attr), + num_nodes=num_nodes, + ) + + +def test_feuler_structure_tokens_are_stable_under_node_relabeling(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + graph = SimpleGraph( + edge_index=[[0, 1, 1, 3], [1, 2, 3, 4]], + x=[9, 1, 4, 7, 4], edge_attr=[3, 2, 5, 6], num_nodes=5) + relabeled = _relabel_graph(graph, {0: 4, 1: 2, 2: 0, 3: 3, 4: 1}) + serializer = FrequencyGuidedEulerianSerializer().fit([graph]) + + original = serializer.serialize(graph) + permuted = serializer.serialize(relabeled) + + assert _feuler_structure_tokens(original.token_ids) == _feuler_structure_tokens( + permuted.token_ids) + + +def test_feuler_frequency_is_stable_under_node_relabeling(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + graph = SimpleGraph( + edge_index=[[0, 1, 1], [1, 2, 3]], + x=[9, 1, 4, 7], edge_attr=[3, 2, 5], num_nodes=4) + relabeled = _relabel_graph(graph, {0: 3, 1: 2, 2: 0, 3: 1}) + + assert FrequencyGuidedEulerianSerializer().fit([graph]).frequency_map == ( + FrequencyGuidedEulerianSerializer().fit([relabeled]).frequency_map) + + +@pytest.mark.parametrize( + "graph, old_to_new", + [ + ( + SimpleGraph( + edge_index=[[0, 1, 2, 3], [1, 2, 3, 0]], + x=[1, 1, 1, 1], edge_attr=[2, 2, 2, 2], num_nodes=4), + {0: 2, 1: 0, 2: 3, 3: 1}, + ), + ( + SimpleGraph( + edge_index=[[0, 0, 0, 0], [1, 2, 3, 4]], + x=[9, 1, 1, 1, 1], edge_attr=[2, 2, 2, 2], num_nodes=5), + {0: 3, 1: 0, 2: 4, 3: 1, 4: 2}, + ), + ( + SimpleGraph( + edge_index=[ + [0, 0, 0, 0, 1, 1, 1, 2, 2, 3], + [1, 2, 3, 4, 2, 3, 4, 3, 4, 4], + ], + x=[1, 1, 1, 1, 1], edge_attr=[2] * 10, num_nodes=5), + {0: 4, 1: 2, 2: 0, 3: 3, 4: 1}, + ), + ], + ids=["cycle", "star_with_equivalent_leaves", "regular_graph"], +) +def test_feuler_symmetric_graphs_compare_by_structural_equivalence(graph, old_to_new): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + relabeled = _relabel_graph(graph, old_to_new) + serializer = FrequencyGuidedEulerianSerializer().fit([graph]) + + original = serializer.serialize(graph) + permuted = serializer.serialize(relabeled) + + assert _feuler_structure_tokens(original.token_ids) == _feuler_structure_tokens( + permuted.token_ids) + _assert_feuler_round_trip(graph) + _assert_feuler_round_trip(relabeled) + + +def _random_node_permutations(num_nodes, seed=0, count=50): + rng = random.Random(seed) + for _ in range(count): + node_ids = list(range(num_nodes)) + rng.shuffle(node_ids) + yield dict(enumerate(node_ids)) + + +def _assert_feuler_permutation_regression(graph, seed=0): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + serializer = FrequencyGuidedEulerianSerializer().fit([graph]) + expected = _feuler_structure_tokens(serializer.serialize(graph).token_ids) + _assert_feuler_round_trip(graph) + for old_to_new in _random_node_permutations(graph.num_nodes, seed=seed): + permuted = _relabel_graph(graph, old_to_new) + assert _feuler_structure_tokens(serializer.serialize(permuted).token_ids) == expected + _assert_feuler_round_trip(permuted) + + +def test_feuler_path_permutation(): + _assert_feuler_permutation_regression(SimpleGraph( + edge_index=[[0, 1, 2, 3], [1, 2, 3, 4]], + x=[8, 1, 4, 7, 9], edge_attr=[3, 2, 5, 6], num_nodes=5), seed=11) + + +def test_feuler_tree_permutation(): + _assert_feuler_permutation_regression(SimpleGraph( + edge_index=[[0, 0, 1, 1, 2], [1, 2, 3, 4, 5]], + x=[9, 1, 4, 7, 3, 8], edge_attr=[1, 2, 3, 4, 5], num_nodes=6), seed=13) + + +def test_feuler_star_permutation(): + _assert_feuler_permutation_regression(SimpleGraph( + edge_index=[[0, 0, 0, 0, 0], [1, 2, 3, 4, 5]], + x=[9, 1, 1, 1, 1, 1], edge_attr=[2, 2, 2, 2, 2], num_nodes=6), seed=17) + + +def test_feuler_cycle_permutation(): + _assert_feuler_permutation_regression(SimpleGraph( + edge_index=[[0, 1, 2, 3, 4], [1, 2, 3, 4, 0]], + x=[1, 1, 1, 1, 1], edge_attr=[2, 2, 2, 2, 2], num_nodes=5), seed=19) + + +def test_feuler_disconnected_permutation(): + _assert_feuler_permutation_regression(SimpleGraph( + edge_index=[[0, 1, 3, 3], [1, 2, 4, 5]], + x=[8, 1, 4, 9, 2, 7], edge_attr=[3, 2, 5, 6], num_nodes=6), seed=23) + + +def test_feuler_frequency_tie_independent_of_node_ids(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + graph = SimpleGraph( + edge_index=[[0, 0, 0], [1, 2, 3]], + x=[0, 7, 8, 9], edge_attr=[1, 2, 3], num_nodes=4) + serializer = FrequencyGuidedEulerianSerializer().fit([graph]) + graph_view = serializer._read_graph(graph) + priorities = [serializer._arc_priority(graph_view, 0, dst, label) + for dst, label in zip([1, 2, 3], [1, 2, 3])] + + assert priorities == [1, 1, 1] + _assert_feuler_permutation_regression(graph, seed=29) + + +def test_feuler_round_trip_after_random_permutation(): + _assert_feuler_permutation_regression(SimpleGraph( + edge_index=[[0, 0, 1, 2, 4], [1, 2, 3, 4, 5]], + x=[9, 1, 4, 7, 3, 8], edge_attr=[1, 2, 3, 4, 5], num_nodes=6), seed=31) + + +def test_round_trip_after_bpe_encode_decode(): + graph = SimpleGraph( + edge_index=[[0, 1, 2], [1, 2, 0]], x=[3, 3, 3], edge_attr=[4, 4, 4]) + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + tokenizer = graph_tokenizer.GraphTokenizer( + bpe=graph_bpe.GraphBPE(num_merges=2, min_frequency=2)).fit([graph, graph]) + + reconstructed = tokenizer.decode_graph(tokenizer.encode_graph(graph)) + + assert reconstructed["num_nodes"] == graph.num_nodes + assert reconstructed["x"] == graph.x + assert _canonical_labeled_edges( + reconstructed["edge_index"], reconstructed["edge_attr"]) == _canonical_labeled_edges( + graph.edge_index, graph.edge_attr) + assert _component_sets( + reconstructed["num_nodes"], reconstructed["edge_index"]) == _component_sets( + graph.num_nodes, graph.edge_index) + + +def test_feuler_rejects_graphs_that_cannot_be_an_undirected_simple_graph(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + graph = SimpleGraph(edge_index=[[0, 1], [1, 0]], x=[4, 5], edge_attr=[1, 2]) + + with pytest.raises(ValueError, match="undirected simple"): + FrequencyGuidedEulerianSerializer().fit([graph]) + + +def test_feuler_rejects_parallel_edges_instead_of_deduplicating_them(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + graph = SimpleGraph(edge_index=[[0, 0], [1, 1]], x=[4, 5], edge_attr=[1, 1]) + + with pytest.raises(ValueError, match="Parallel edges"): + FrequencyGuidedEulerianSerializer().fit([graph]) + + +def test_feuler_rejects_self_loops_and_explicit_directed_graphs(): + FrequencyGuidedEulerianSerializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer + self_loop = SimpleGraph(edge_index=[[0], [0]], x=[4], edge_attr=[1]) + directed = { + "edge_index": [[0], [1]], "x": [4, 5], "edge_attr": [1], + "directed": True, + } + + with pytest.raises(ValueError, match="self-loops"): + FrequencyGuidedEulerianSerializer().serialize(self_loop) + with pytest.raises(ValueError, match="directed graph semantics"): + FrequencyGuidedEulerianSerializer().serialize(directed) + + +def test_serializer_preserves_preencoded_paper_node_and_edge_tokens(): + graph_serializer = _load_graph_serializer_module() + graph = SimpleGraph(edge_index=[[0], [1]], x=[13, 17], edge_attr=[2]) + + result = graph_serializer.FrequencyGuidedEulerianSerializer().fit([graph]).serialize(graph) + + serializer = graph_serializer.FrequencyGuidedEulerianSerializer().fit([graph]) + assert result.token_ids == [ + serializer.node_token(13), serializer.node_reference_token(0), + serializer.edge_token(2), + serializer.node_token(17), serializer.node_reference_token(1), + serializer.edge_token(2), + serializer.node_token(13), serializer.node_reference_token(0), + ] + + +def test_eulerian_serializer_does_not_collect_frequency_statistics(): + graph_serializer = _load_graph_serializer_module() + graph = SimpleGraph(edge_index=[[0], [1]], x=[13, 17], edge_attr=[2]) + + serializer = graph_serializer.EulerianSerializer().fit([graph]) + + assert serializer.name == "eulerian" + assert serializer.frequency_map == {} + assert serializer.serialize(graph).metadata["method"] == "eulerian" + + +def test_graph_bpe_trains_most_frequent_pairs_and_encodes_sequences(): + GraphBPE = _load_graph_bpe_module().GraphBPE + + bpe = GraphBPE(num_merges=2, min_frequency=2) + bpe.fit([[1, 2, 1, 2], [1, 2, 1, 2], [1, 2, 3], [4, 5]]) + + assert bpe.codebook.merge_rules[0] == (1, 2, 6) + assert bpe.codebook.merge_rules[1] == (6, 6, 7) + assert bpe.codebook.vocab_size == 8 + assert bpe.encode([1, 2, 1, 2, 3]) == [7, 3] + assert bpe.batch_encode([[1, 2], [4, 5]]) == [[6], [4, 5]] + + +def test_graph_bpe_never_merges_pairs_touching_protected_tokens(): + GraphBPE = _load_graph_bpe_module().GraphBPE + + bpe = GraphBPE(num_merges=2, min_frequency=2, protected_token_ids={7}) + bpe.fit([[10, 7, 11], [10, 7, 11]]) + + assert bpe.codebook.merge_rules == [] + assert bpe.encode([10, 7, 11]) == [10, 7, 11] + + +def test_graph_bpe_cpp_trains_protected_segments_with_the_native_backend(monkeypatch): + GraphBPE = _load_graph_bpe_module().GraphBPE + from third_party import graph_bpe_cpp + + calls = [] + + monkeypatch.setattr(graph_bpe_cpp, "is_available", lambda: True) + monkeypatch.setattr( + graph_bpe_cpp, + "train_bpe", + lambda sequences, num_merges, min_frequency, initial_vocab_size: calls.append( + (sequences, num_merges, min_frequency, initial_vocab_size)) or { + "merge_rules": [(1, 2, 8)], + "vocab_size": 9, + "metadata": {"backend": "cpp"}, + }, + ) + + bpe = GraphBPE( + num_merges=1, + min_frequency=2, + backend="cpp", + protected_token_ids={7}, + ).fit([[1, 2, 7, 3, 4], [1, 2, 7, 3, 4]]) + + assert calls == [([[1, 2], [3, 4], [1, 2], [3, 4]], 1, 2, 8)] + assert bpe.codebook.metadata["backend"] == "cpp" + + +def test_graph_bpe_cpp_rejects_missing_native_backend_with_protected_tokens(monkeypatch): + GraphBPE = _load_graph_bpe_module().GraphBPE + from third_party import graph_bpe_cpp + + monkeypatch.setattr(graph_bpe_cpp, "is_available", lambda: False) + + with pytest.raises(ImportError, match="native extension"): + GraphBPE(backend="cpp", protected_token_ids={7}).fit([[1, 7, 2]]) + + +def test_graph_bpe_respects_min_frequency_and_roundtrips_codebook(tmp_path): + graph_bpe = _load_graph_bpe_module() + GraphBPE = graph_bpe.GraphBPE + + bpe = GraphBPE(num_merges=5, min_frequency=3).fit([[9, 10], [9, 10], [1, 2]]) + + assert bpe.codebook.merge_rules == [] + assert bpe.encode([9, 10]) == [9, 10] + + trained = GraphBPE(num_merges=1, min_frequency=2).fit([[9, 10], [9, 10], [1, 2]]) + codebook_path = tmp_path / "codebook.json" + trained.save_codebook(codebook_path) + + loaded = GraphBPE.load_codebook(codebook_path) + assert loaded.codebook.merge_rules == [(9, 10, 11)] + assert loaded.encode([9, 10, 9, 10]) == [11, 11] + + +def test_graph_bpe_auto_backend_matches_python_fallback(): + GraphBPE = _load_graph_bpe_module().GraphBPE + + sequences = [[1, 2, 1, 2], [1, 2, 1, 2], [1, 2, 3], [4, 5]] + python_bpe = GraphBPE(num_merges=2, min_frequency=2, backend="python").fit(sequences) + auto_bpe = GraphBPE(num_merges=2, min_frequency=2, backend="auto").fit(sequences) + + assert auto_bpe.codebook.merge_rules == python_bpe.codebook.merge_rules + assert auto_bpe.codebook.vocab_size == python_bpe.codebook.vocab_size + assert auto_bpe.codebook.metadata["backend"] in {"cpp", "python"} + assert auto_bpe.encode([1, 2, 1, 2, 3]) == python_bpe.encode([1, 2, 1, 2, 3]) + + +def test_graph_bpe_cpp_matches_python_with_protected_separators_when_available(): + GraphBPE = _load_graph_bpe_module().GraphBPE + from third_party import graph_bpe_cpp + + if not graph_bpe_cpp.is_available(): + pytest.skip("graph_bpe_cpp native extension is unavailable") + separator = 7 + sequences = [[1, 2, separator, 1, 2], [1, 2, separator, 1, 2], [1, 2]] + python_bpe = GraphBPE( + num_merges=2, min_frequency=2, backend="python", + protected_token_ids={separator}).fit(sequences) + cpp_bpe = GraphBPE( + num_merges=2, min_frequency=2, backend="cpp", + protected_token_ids={separator}).fit(sequences) + + assert cpp_bpe.codebook.merge_rules == python_bpe.codebook.merge_rules + assert cpp_bpe.codebook.vocab_size == python_bpe.codebook.vocab_size + assert cpp_bpe.encode(sequences[0]) == python_bpe.encode(sequences[0]) + assert separator in cpp_bpe.encode(sequences[0]) + + +def test_graph_bpe_auto_falls_back_without_the_native_bridge(monkeypatch): + GraphBPE = _load_graph_bpe_module().GraphBPE + from third_party import graph_bpe_cpp + + monkeypatch.setattr(graph_bpe_cpp, "is_available", lambda: False) + monkeypatch.setattr( + graph_bpe_cpp, + "train_bpe", + lambda *_args, **_kwargs: (_ for _ in ()).throw(AssertionError("native bridge used")), + ) + + bpe = GraphBPE(num_merges=1, min_frequency=2, backend="auto").fit( + [[1, 2], [1, 2]]) + + assert bpe.codebook.merge_rules == [(1, 2, 3)] + assert bpe.codebook.metadata["backend"] == "python" + + +def test_graph_bpe_cpp_requires_the_native_bridge(monkeypatch): + GraphBPE = _load_graph_bpe_module().GraphBPE + from third_party import graph_bpe_cpp + + monkeypatch.setattr(graph_bpe_cpp, "is_available", lambda: False) + + bpe = GraphBPE(num_merges=1, min_frequency=2, backend="cpp") + with pytest.raises(ImportError, match="native extension"): + bpe.fit([[1, 2], [1, 2]]) + + bpe.codebook.merge_rules = [(1, 2, 3)] + with pytest.raises(ImportError, match="native extension"): + bpe.encode([1, 2]) + with pytest.raises(ImportError, match="native extension"): + bpe.batch_encode([[1, 2]]) + + +def test_graph_bpe_cpp_package_exposes_native_results_when_available(): + from third_party import graph_bpe_cpp + + if not graph_bpe_cpp.is_available(): + with pytest.raises(ImportError, match="native extension"): + graph_bpe_cpp.train_bpe([[1, 2]], num_merges=1, min_frequency=2) + return + + result = graph_bpe_cpp.train_bpe( + [[1, 2, 1, 2], [1, 2, 1, 2], [1, 2, 3]], + num_merges=2, + min_frequency=2, + ) + + assert result["merge_rules"] == [(1, 2, 4), (4, 4, 5)] + assert result["vocab_size"] == 6 + assert graph_bpe_cpp.encode([1, 2, 1, 2, 3], result["merge_rules"]) == [5, 3] + + +def test_graph_bpe_cpp_batch_encode_crosses_the_bridge_once(monkeypatch): + graph_bpe = _load_graph_bpe_module() + from third_party import graph_bpe_cpp + + calls = [] + + def batch_encode(token_sequences, merge_rules): + calls.append((token_sequences, merge_rules)) + return [[6], [4, 5]] + + def fail_single_encode(*_args, **_kwargs): + raise AssertionError("batch encoding must not cross the native bridge once per sequence") + + monkeypatch.setattr(graph_bpe_cpp, "batch_encode", batch_encode, raising=False) + monkeypatch.setattr(graph_bpe_cpp, "encode", fail_single_encode) + monkeypatch.setattr(graph_bpe_cpp, "is_available", lambda: True) + engine = graph_bpe.GraphBPE(num_merges=0, backend="cpp") + engine.codebook = graph_bpe.BPECodebook(merge_rules=[(1, 2, 6)], vocab_size=7) + + encoded = engine.batch_encode([[1, 2], [4, 5]]) + + assert encoded == [[6], [4, 5]] + assert calls == [([[1, 2], [4, 5]], [(1, 2, 6)])] + + +def test_graph_bpe_auto_uses_native_batch_for_native_codebook(monkeypatch): + graph_bpe = _load_graph_bpe_module() + from third_party import graph_bpe_cpp + + calls = [] + monkeypatch.setattr( + graph_bpe_cpp, + "batch_encode", + lambda sequences, rules: calls.append((sequences, rules)) or [[6]], + ) + monkeypatch.setattr(graph_bpe_cpp, "is_available", lambda: True) + engine = graph_bpe.GraphBPE(num_merges=0, backend="auto") + engine.codebook = graph_bpe.BPECodebook( + merge_rules=[(1, 2, 6)], + vocab_size=7, + metadata={"backend": "cpp"}, + ) + + assert engine.batch_encode([[1, 2]]) == [[6]] + assert calls == [([[1, 2]], [(1, 2, 6)])] + + +def test_graph_tokenizer_fits_serializer_and_bpe_then_encodes_graph(): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + GraphTokenizer = graph_tokenizer.GraphTokenizer + GraphBPE = graph_bpe.GraphBPE + + graph = SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]) + tokenizer = GraphTokenizer(bpe=GraphBPE(num_merges=1, min_frequency=2)) + tokenizer.fit([graph, graph]) + + encoding = tokenizer.encode_graph(graph) + + assert tokenizer.bpe.codebook.merge_rules + assert encoding.input_ids[0] == tokenizer.special_tokens.cls_token_id + assert encoding.input_ids[-1] == tokenizer.special_tokens.sep_token_id + assert encoding.attention_mask == [1] * len(encoding.input_ids) + assert encoding.metadata["serializer"]["method"] == "feuler" + + +def test_graph_tokenizer_protects_component_separator_from_bpe_merges(): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + graph = SimpleGraph( + edge_index=[[0, 2], [1, 3]], x=[1, 2, 1, 2], edge_attr=[5, 5]) + tokenizer = graph_tokenizer.GraphTokenizer( + bpe=graph_bpe.GraphBPE(num_merges=8, min_frequency=2)) + + tokenizer.fit([graph, graph]) + + separator = tokenizer.special_tokens.component_sep_token_id + assert all(separator not in rule[:2] for rule in tokenizer.bpe.codebook.merge_rules) + + +def test_graph_tokenizer_rejects_empty_training_corpus_and_unfitted_encoding(): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + graph = SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]) + tokenizer = graph_tokenizer.GraphTokenizer(bpe=graph_bpe.GraphBPE(num_merges=0)) + + with pytest.raises(ValueError, match="empty training corpus"): + tokenizer.fit([]) + with pytest.raises(RuntimeError, match="must be fit"): + tokenizer.encode_graph(graph) + + +def test_graph_tokenizer_records_train_only_fit_provenance(): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + graph = SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]) + tokenizer = graph_tokenizer.GraphTokenizer(bpe=graph_bpe.GraphBPE(num_merges=0)) + + tokenizer.fit([graph], graph_ids=[42]) + encoding = tokenizer.encode_graph(graph) + + assert tokenizer.fit_graph_ids_hash + assert encoding.metadata["tokenizer_schema_version"] == 1 + assert encoding.metadata["fit_graph_ids_hash"] == tokenizer.fit_graph_ids_hash + + +def test_frozen_train_vocabulary_maps_unseen_validation_and_test_tokens_to_unk(): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + train = SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]) + held_out = SimpleGraph(edge_index=[[0], [1]], x=[101, 102], edge_attr=[5]) + tokenizer = graph_tokenizer.GraphTokenizer( + bpe=graph_bpe.GraphBPE(num_merges=1, min_frequency=2)).fit([train, train]) + before = (dict(tokenizer.vocabulary), list(tokenizer.bpe.codebook.merge_rules)) + + validation = tokenizer.encode_graph(held_out) + test = tokenizer.encode_graph(held_out) + + assert (dict(tokenizer.vocabulary), list(tokenizer.bpe.codebook.merge_rules)) == before + assert tokenizer.special_tokens.unk_token_id in validation.input_ids + assert tokenizer.special_tokens.unk_token_id in test.input_ids + assert all(0 <= token <= tokenizer.max_token_id for token in validation.input_ids) + assert all(0 <= token <= tokenizer.max_token_id for token in test.input_ids) + + +def test_train_vocabulary_and_bpe_merge_ids_are_contiguous_and_unique(): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + graph = SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]) + tokenizer = graph_tokenizer.GraphTokenizer( + bpe=graph_bpe.GraphBPE(num_merges=2, min_frequency=2)).fit([graph, graph]) + vocabulary_ids = sorted(tokenizer.vocabulary.values()) + merge_ids = [rule[2] for rule in tokenizer.bpe.codebook.merge_rules] + + assert vocabulary_ids == list(range(8, tokenizer.max_token_id + 1)) + assert len(merge_ids) == len(set(merge_ids)) + + +def test_graph_tokenizer_batch_encodes_serialized_graphs_together(monkeypatch): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + graph_one = SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]) + graph_two = SimpleGraph(edge_index=[[0], [1]], x=[1, 3], edge_attr=[6]) + tokenizer = graph_tokenizer.GraphTokenizer( + bpe=graph_bpe.GraphBPE(num_merges=0)).fit([graph_one, graph_two]) + calls = [] + + def batch_encode(sequences): + calls.append(sequences) + return [list(sequence) for sequence in sequences] + + monkeypatch.setattr(tokenizer.bpe, "batch_encode", batch_encode) + + results = tokenizer.batch_encode_graphs([graph_one, graph_two]) + + assert len(calls) == 1 + assert len(calls[0]) == 2 + assert [result.input_ids[0] for result in results] == [ + tokenizer.special_tokens.cls_token_id, + tokenizer.special_tokens.cls_token_id, + ] + + +def test_graph_tokenizer_call_attaches_token_fields_to_graph(): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + GraphTokenizer = graph_tokenizer.GraphTokenizer + GraphBPE = graph_bpe.GraphBPE + + graph = SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]) + tokenizer = GraphTokenizer(bpe=GraphBPE(num_merges=1, min_frequency=2)).fit([graph, graph]) + + transformed = tokenizer(graph) + + assert transformed is graph + assert graph.input_ids == tokenizer.encode_graph(graph).input_ids + assert graph.attention_mask == [1] * len(graph.input_ids) + assert graph.graph_tokenizer_metadata["serializer"]["method"] == "feuler" + + +def test_graph_tokenizer_builds_mlm_pretraining_batch(): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + GraphTokenizer = graph_tokenizer.GraphTokenizer + GraphBPE = graph_bpe.GraphBPE + + graph = SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]) + tokenizer = GraphTokenizer(bpe=GraphBPE(num_merges=0, min_frequency=2)).fit([graph]) + + batch = tokenizer.build_mlm_batch([graph], max_length=10, mask_prob=1.0, seed=3) + sequence = batch.input_ids[0] + labels = batch.labels[0] + + assert len(sequence) == 10 + assert len(batch.attention_mask[0]) == 10 + assert sequence[0] == tokenizer.special_tokens.cls_token_id + assert labels[0] == -100 + assert tokenizer.special_tokens.mask_token_id in sequence + assert sum(1 for value in labels if value != -100) > 0 + assert batch.metadata["num_mlm_labels"] == sum(1 for value in labels if value != -100) + + +def test_graph_tokenizer_builds_mlm_batch_from_token_sequences(): + tokenizer = _load_graph_tokenizer_module().GraphTokenizer() + + batch = tokenizer.build_mlm_batch_from_token_sequences( + [[3, 6, 4], [3, 8, 9, 4]], max_length=5, mask_prob=1.0, seed=7) + + assert batch.input_ids == [[3, 2, 4, 0, 0], [3, 2, 2, 4, 0]] + assert batch.attention_mask == [[1, 1, 1, 0, 0], [1, 1, 1, 1, 0]] + assert batch.labels == [[-100, 6, -100, -100, -100], [-100, 8, 9, -100, -100]] + + +def test_graph_tokenizer_rejects_sequence_overflow_instead_of_truncating(): + tokenizer = _load_graph_tokenizer_module().GraphTokenizer() + + with pytest.raises(ValueError, match="exceeds max_length"): + tokenizer.build_mlm_batch_from_token_sequences( + [[3, 6, 7, 4]], max_length=3, mask_prob=0.0) + + +def test_serializer_rejects_multidimensional_features_without_an_adapter(): + serializer = _load_graph_serializer_module().FrequencyGuidedEulerianSerializer() + graph = SimpleGraph(edge_index=[[0], [1]], x=[[1, 2], [1, 9]], edge_attr=[[5, 6]]) + + with pytest.raises(ValueError, match="Multi-dimensional"): + serializer.fit([graph]) + + +def test_graph_tokenizer_mlm_forces_one_mask_when_random_draw_selects_none(): + graph_tokenizer = _load_graph_tokenizer_module() + tokenizer = graph_tokenizer.GraphTokenizer() + input_ids = [[ + tokenizer.special_tokens.cls_token_id, + 11, + 12, + tokenizer.special_tokens.sep_token_id, + ]] + + masked, labels = tokenizer._mask_input_ids( + input_ids, mask_prob=1e-12, seed=0) + + selected = [ + (row, column) + for row, label_row in enumerate(labels) + for column, label in enumerate(label_row) + if label != -100 + ] + assert len(selected) == 1 + row, column = selected[0] + assert labels[row][column] == input_ids[row][column] + assert masked[row][column] == tokenizer.special_tokens.mask_token_id + + +def test_graph_tokenizer_replaces_component_separator_before_mlm(): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + GraphTokenizer = graph_tokenizer.GraphTokenizer + GraphBPE = graph_bpe.GraphBPE + + graph = SimpleGraph(edge_index=[[0, 2], [1, 3]], x=[1, 2, 3, 4], edge_attr=[5, 6]) + tokenizer = GraphTokenizer(bpe=GraphBPE(num_merges=0, min_frequency=2)).fit([graph]) + + encoding = tokenizer.encode_graph(graph) + batch = tokenizer.build_mlm_batch([graph], max_length=20, mask_prob=1.0, seed=4) + sep_index = encoding.input_ids.index(tokenizer.special_tokens.component_sep_token_id) + + assert -1 not in encoding.input_ids + assert tokenizer.special_tokens.component_sep_token_id in encoding.input_ids + assert batch.input_ids[0][sep_index] == tokenizer.special_tokens.component_sep_token_id + assert batch.labels[0][sep_index] == -100 + + +def test_graph_tokenizer_rejects_a_model_vocabulary_smaller_than_encoded_tokens(): + graph_tokenizer = _load_graph_tokenizer_module() + graph_bpe = _load_graph_bpe_module() + GraphTokenizer = graph_tokenizer.GraphTokenizer + GraphBPE = graph_bpe.GraphBPE + + graph = SimpleGraph(edge_index=[[0], [1]], x=[1, 2], edge_attr=[5]) + tokenizer = GraphTokenizer(bpe=GraphBPE(num_merges=1, min_frequency=2)).fit([graph, graph]) + + with pytest.raises(ValueError, match="vocab_size"): + tokenizer.validate_model_vocab(2) diff --git a/third_party/__init__.py b/third_party/__init__.py new file mode 100644 index 000000000..7a9a5e642 --- /dev/null +++ b/third_party/__init__.py @@ -0,0 +1 @@ +"""Optional native extensions shipped with GammaGL.""" diff --git a/third_party/graph_bpe_cpp/__init__.py b/third_party/graph_bpe_cpp/__init__.py new file mode 100644 index 000000000..f1332b2fc --- /dev/null +++ b/third_party/graph_bpe_cpp/__init__.py @@ -0,0 +1,104 @@ +from typing import List, Sequence, Tuple + + +MergeRule = Tuple[int, int, int] + + +def _load_native(): + try: + from . import _graph_bpe + except ImportError: + return None + return _graph_bpe + + +def is_available() -> bool: + """Return whether the optional native BPE extension is importable.""" + + return _load_native() is not None + + +def backend_name() -> str: + return "cpp" if is_available() else "python" + + +def train_bpe(token_sequences, num_merges: int, min_frequency: int, initial_vocab_size=None): + native = _load_native() + if native is None: + raise ImportError("graph_bpe_cpp requires the optional native extension.") + return _normalize_result( + native.train_bpe( + token_sequences, + int(num_merges), + int(min_frequency), + -1 if initial_vocab_size is None else int(initial_vocab_size), + ), "cpp") + + +def encode(token_sequence, merge_rules: Sequence[Sequence[int]]) -> List[int]: + native = _load_native() + if native is None: + raise ImportError("graph_bpe_cpp requires the optional native extension.") + return [int(token) for token in native.encode(token_sequence, merge_rules)] + + +def batch_encode(token_sequences, merge_rules: Sequence[Sequence[int]]) -> List[List[int]]: + native = _load_native() + if native is None: + raise ImportError("graph_bpe_cpp requires the optional native extension.") + return [ + [int(token) for token in sequence] + for sequence in native.batch_encode(token_sequences, merge_rules) + ] + + +def _train_bpe_python(token_sequences, num_merges: int, min_frequency: int): + """Reference implementation used to verify the optional native backend.""" + sequences = [[int(token) for token in sequence] for sequence in token_sequences] + next_token_id = max((token for sequence in sequences for token in sequence), default=-1) + 1 + merge_rules = [] + for _ in range(int(num_merges)): + counts = {} + for sequence in sequences: + for pair in zip(sequence, sequence[1:]): + counts[pair] = counts.get(pair, 0) + 1 + if not counts: + break + pair = min(counts, key=lambda item: (-counts[item], item)) + if counts[pair] < int(min_frequency): + break + rule = (*pair, next_token_id) + merge_rules.append(rule) + sequences = [_apply_merge(sequence, rule) for sequence in sequences] + next_token_id += 1 + return {"merge_rules": merge_rules, "vocab_size": next_token_id} + + +def _apply_merge(sequence, rule): + left, right, merged = rule + encoded = [] + index = 0 + while index < len(sequence): + if index + 1 < len(sequence) and sequence[index:index + 2] == [left, right]: + encoded.append(merged) + index += 2 + else: + encoded.append(sequence[index]) + index += 1 + return encoded + + +def _normalize_result(result, backend: str): + merge_rules = [_as_merge_rule(rule) for rule in result.get("merge_rules", [])] + return { + "merge_rules": merge_rules, + "vocab_size": int(result.get("vocab_size", 0)), + "metadata": { + **{key: int(value) for key, value in result.get("metadata", {}).items() if isinstance(value, int)}, + "backend": backend, + }, + } + + +def _as_merge_rule(rule) -> MergeRule: + return tuple(int(value) for value in rule) diff --git a/third_party/graph_bpe_cpp/_graph_bpe.cpp b/third_party/graph_bpe_cpp/_graph_bpe.cpp new file mode 100644 index 000000000..83c12e9af --- /dev/null +++ b/third_party/graph_bpe_cpp/_graph_bpe.cpp @@ -0,0 +1,117 @@ +#include +#include + +#include +#include +#include +#include + +namespace py = pybind11; + +using Sequence = std::vector; +using MergeRule = std::tuple; + +long long next_token_id(const std::vector& sequences) { + long long max_token = -1; + for (const auto& sequence : sequences) { + for (long long token : sequence) { + max_token = std::max(max_token, token); + } + } + return max_token + 1; +} + +std::map, long long> count_pairs(const std::vector& sequences) { + std::map, long long> counts; + for (const auto& sequence : sequences) { + for (std::size_t i = 0; i + 1 < sequence.size(); ++i) { + counts[{sequence[i], sequence[i + 1]}] += 1; + } + } + return counts; +} + +Sequence apply_merge(const Sequence& sequence, const MergeRule& rule) { + auto [left, right, new_id] = rule; + Sequence merged; + for (std::size_t i = 0; i < sequence.size();) { + if (i + 1 < sequence.size() && sequence[i] == left && sequence[i + 1] == right) { + merged.push_back(new_id); + i += 2; + } else { + merged.push_back(sequence[i]); + i += 1; + } + } + return merged; +} + +py::dict train_bpe(const std::vector& input_sequences, int num_merges, int min_frequency, + long long initial_vocab_size) { + std::vector sequences = input_sequences; + std::vector merge_rules; + long long next_id = std::max(next_token_id(sequences), initial_vocab_size); + + for (int merge_index = 0; merge_index < num_merges; ++merge_index) { + auto counts = count_pairs(sequences); + if (counts.empty()) { + break; + } + + auto best = counts.begin(); + for (auto it = counts.begin(); it != counts.end(); ++it) { + if (it->second > best->second || + (it->second == best->second && it->first < best->first)) { + best = it; + } + } + if (best->second < min_frequency) { + break; + } + + MergeRule rule{best->first.first, best->first.second, next_id}; + merge_rules.push_back(rule); + for (auto& sequence : sequences) { + sequence = apply_merge(sequence, rule); + } + next_id += 1; + } + + py::dict metadata; + metadata["num_merges_requested"] = num_merges; + metadata["num_merges_performed"] = static_cast(merge_rules.size()); + metadata["min_frequency"] = min_frequency; + + py::dict result; + result["merge_rules"] = merge_rules; + result["vocab_size"] = next_id; + result["metadata"] = metadata; + return result; +} + +Sequence encode(const Sequence& sequence, const std::vector& merge_rules) { + Sequence encoded = sequence; + for (const auto& rule : merge_rules) { + encoded = apply_merge(encoded, rule); + } + return encoded; +} + +std::vector batch_encode( + const std::vector& sequences, + const std::vector& merge_rules) { + std::vector encoded; + encoded.reserve(sequences.size()); + for (const auto& sequence : sequences) { + encoded.push_back(encode(sequence, merge_rules)); + } + return encoded; +} + +PYBIND11_MODULE(_graph_bpe, m) { + m.doc() = "Native GraphTokenizer BPE backend"; + m.def("train_bpe", &train_bpe, py::arg("token_sequences"), py::arg("num_merges"), + py::arg("min_frequency"), py::arg("initial_vocab_size") = -1); + m.def("encode", &encode, py::arg("token_sequence"), py::arg("merge_rules")); + m.def("batch_encode", &batch_encode, py::arg("token_sequences"), py::arg("merge_rules")); +} diff --git a/third_party/graph_bpe_cpp/setup.py b/third_party/graph_bpe_cpp/setup.py new file mode 100644 index 000000000..057c0ff29 --- /dev/null +++ b/third_party/graph_bpe_cpp/setup.py @@ -0,0 +1,30 @@ +from pathlib import Path + +from setuptools import Extension, setup + + +try: + import pybind11 +except ImportError as exc: + raise SystemExit("pybind11 is required to build graph_bpe_cpp. Install pybind11 first.") from exc + + +HERE = Path(__file__).resolve().parent + + +setup( + name="graph_bpe_cpp", + version="0.1.0", + description="Optional native GraphTokenizer BPE backend.", + packages=["third_party.graph_bpe_cpp"], + package_dir={"third_party.graph_bpe_cpp": str(HERE)}, + ext_modules=[ + Extension( + "third_party.graph_bpe_cpp._graph_bpe", + sources=[str(HERE / "_graph_bpe.cpp")], + include_dirs=[pybind11.get_include()], + language="c++", + extra_compile_args=["/std:c++17"] if __import__("os").name == "nt" else ["-std=c++17"], + ) + ], +) From a0d80044a7a7ebad55aaed63cc578a0bc087c1ec Mon Sep 17 00:00:00 2001 From: baihchou8787 Date: Mon, 31 Aug 2026 02:08:29 +0000 Subject: [PATCH 2/2] docs(graph-tokenizer): update Chinese paper guide --- examples/graph_tokenizer/README.md | 110 ++++++++++++++--------------- 1 file changed, 52 insertions(+), 58 deletions(-) diff --git a/examples/graph_tokenizer/README.md b/examples/graph_tokenizer/README.md index 221aef93d..18dd59596 100644 --- a/examples/graph_tokenizer/README.md +++ b/examples/graph_tokenizer/README.md @@ -1,73 +1,67 @@ # GraphTokenizer -This example provides graph serialization, Graph BPE tokenization, masked-language-model pretraining, supervised fine-tuning, checkpointing, and multi-seed evaluation for GraphTokenizer. +本示例为 GraphTokenizer 提供图序列化、Graph BPE 词元化、掩码语言模型预训练、监督微调、检查点保存和多随机种子评估功能。 -## Paper +## 论文 -GraphTokenizer is described in [Graph Tokenization for Bridging Graphs and Transformers](https://openreview.net/forum?id=jCctxI1BGF) (ICLR 2026). +GraphTokenizer 由论文[连接图与 Transformer 的图词元化](https://openreview.net/forum?id=jCctxI1BGF)(ICLR 2026)提出。 -The paper protocol uses frequency-guided Eulerian serialization (Feuler), fits Graph BPE on the training split only, pretrains with masked language modeling, selects the best fine-tuning checkpoint by validation performance, evaluates the test split once, and reports the mean and population standard deviation over five runs. +论文实验协议采用频率引导的欧拉序列化(Feuler),仅在训练集划分上拟合 Graph BPE,使用掩码语言模型进行预训练,根据验证集性能选择最佳微调检查点,只在测试集上评估一次,并报告五次运行结果的均值和总体标准差。 -The Feuler serializer is reversible for its supported simple-undirected graph -domain. Paper GTE runs load the pinned official encoder into the native TLX -GraphGTE implementation; graph-token embeddings and task heads are new -initializations. +Feuler 序列化器在其支持的无向简单图范围内是可逆的。论文中的 GTE +实验会将版本固定的官方编码器加载到原生 TLX GraphGTE 实现中;图词元嵌入和 +任务头均采用全新初始化。 -Feuler accepts only simple undirected graphs: self-loops and parallel edges -are rejected, and graph objects that explicitly declare `directed=True` or -`is_directed()=True` are rejected. COO may store each undirected edge once or -in symmetric form; both are canonicalized to the same undirected edge. A raw -single-direction COO has no way to encode a distinct directed-graph meaning, -so it is interpreted only as that supported undirected storage form. +Feuler 仅接受无向简单图:不支持自环和平行边,也不接受显式声明 +`directed=True` 或 `is_directed()=True` 的图对象。COO 可以只存储每条无向边 +一次,也可以采用对称形式存储;两种表示都会被规范化为相同的无向边。原始的 +单向 COO 无法表达独立的有向图语义,因此只会被解释为受支持的无向图存储形式。 -## Datasets +## 数据集 -The paper commands below cover: +以下论文实验命令覆盖三个数据集: -- QM9: joint 16-target regression, reported with MAE. -- OGBG-molhiv: binary classification, reported with OGB ROC-AUC. -- Peptides-struct: 11-target regression, reported with Average MAE. +- QM9:联合 16 目标回归,使用 MAE 报告结果。 +- OGBG-molhiv:二分类,使用 OGB ROC-AUC 报告结果。 +- Peptides-struct:11 目标回归,使用平均 MAE 报告结果。 -The GammaGL dataset classes download the official GraphTokenizer release bundle -on first use. The downloaded archive is verified before extraction using its -published SHA-256: +GammaGL 数据集类会在首次使用时下载 GraphTokenizer 官方发布的数据包。下载的 +归档文件会在解压前使用官方公布的 SHA-256 进行校验: ``` 5c437c3c0d4b7278379c0e70d57f98148e5c815d753d8cf68e2a45952bcce459 ``` -It is cached under `/.graph_tokenizer_release`, then copied into -each dataset's normal `raw/` directory. A local archive supplied through -`GAMMAGL_GRAPH_TOKENIZER_DATA_BUNDLE=/path/to/bundle` is also verified against -the same digest. A directory value is an explicit local-development override; -it is never downloaded from a user-provided URL. The released `data.pkl(.gz)` -files are deserialized only after the automatic remote bundle has passed this -verification. +数据包会缓存在 `/.graph_tokenizer_release` 下,然后复制到各数据集的 +常规 `raw/` 目录。通过 +`GAMMAGL_GRAPH_TOKENIZER_DATA_BUNDLE=/path/to/bundle` 指定的本地归档文件也会 +使用相同的摘要进行校验。将该变量设置为目录表示显式启用本地开发覆盖;程序 +不会从用户提供的 URL 下载文件。只有自动下载的远程数据包通过校验后,程序才会 +反序列化其中发布的 `data.pkl(.gz)` 文件。 -## Requirements +## 环境要求 -Paper mode requires: +论文模式要求: - `TL_BACKEND=torch` -- PyTorch 2.1.2 with CUDA 12.1 -- DGL 2.4.0 with CUDA 12.1 +- PyTorch 2.1.2,配套 CUDA 12.1 +- DGL 2.4.0,配套 CUDA 12.1 - PyTorch Geometric 2.4.0 -- TensorLayerX, NumPy, OGB, and the GammaGL package -- `huggingface-hub` and `safetensors` for the pinned native-TLX GTE checkpoint converter -- the native Graph BPE extension for the commands below +- TensorLayerX、NumPy、OGB 和 GammaGL 软件包 +- 用于固定版本原生 TLX GTE 检查点转换的 `huggingface-hub` 和 `safetensors` +- 执行以下命令所需的原生 Graph BPE 扩展 -Install these without adding paper-only packages to GammaGL's core dependency -set: +使用以下方式安装依赖,不会将论文专用软件包加入 GammaGL 的核心依赖集合: ```bash pip install -e '.[graph-tokenizer-paper]' -# Or use examples/graph_tokenizer/requirements.txt in the pinned paper environment. +# 也可以在版本固定的论文环境中使用 examples/graph_tokenizer/requirements.txt。 ``` -`transformers` is optional: it is used only by checkpoint/reference-equivalence -tests, never by the formal TLX GraphBERT/GraphGTE training forward path. +`transformers` 是可选依赖:它仅用于检查点和参考实现等价性测试,不会用于正式的 +TLX GraphBERT/GraphGTE 训练前向传播路径。 -Build the native Graph BPE backend from the repository root: +在仓库根目录编译原生 Graph BPE 后端: ```bash export TL_BACKEND=torch @@ -75,11 +69,11 @@ export DGLBACKEND=pytorch python third_party/graph_bpe_cpp/setup.py build_ext --inplace ``` -The paper models use public TensorLayerX layers and operations directly; Hugging Face Transformers is not required by GraphBERT or GraphGTE. +论文模型直接使用 TensorLayerX 的公共层和运算;GraphBERT 和 GraphGTE 不依赖 Hugging Face Transformers。 -## How to Run +## 运行方法 -All paper hyperparameters are passed explicitly through `argparse`. The commands use repository-relative data, cache, and result paths and can be run directly from the GammaGL repository root. +所有论文超参数均通过 `argparse` 显式传入。以下命令使用相对于仓库的数据、缓存和结果路径,可以直接在 GammaGL 仓库根目录运行。 ### QM9 + BERT @@ -273,9 +267,9 @@ python examples/graph_tokenizer/graph_tokenizer_trainer.py \ --output-dir logs/graph_tokenizer/peptides_struct_gte ``` -Before a long run, add `--preflight` to the corresponding command to validate the dataset, model, runtime, and BPE backend without training. `--resume` restores the latest phase, optimizer, scheduler, random-number-generator, and AMP scaler state from `last_state.pt`; `best.pt` remains the lightweight best-model checkpoint. +在开始长时间运行前,可以向相应命令添加 `--preflight`,在不训练的情况下检查数据集、模型、运行环境和 BPE 后端。`--resume` 会从 `last_state.pt` 恢复最近的训练阶段、优化器、调度器、随机数生成器和 AMP 缩放器状态;`best.pt` 仍然是轻量级的最佳模型检查点。 -For a small non-paper smoke test: +执行一个非论文模式的小型冒烟测试: ```bash python examples/graph_tokenizer/graph_tokenizer_trainer.py \ @@ -286,19 +280,19 @@ python examples/graph_tokenizer/graph_tokenizer_trainer.py \ --bpe-backend python ``` -## Results +## 实验结果 -The paper reports the following five-run mean results: +论文报告的五次运行平均结果如下: -| Dataset | Encoder | Metric | Paper | GammaGL status | +| 数据集 | 编码器 | 指标 | 论文结果 | GammaGL 状态 | | --- | --- | --- | ---: | --- | -| QM9 | BERT | raw MAE ↓ | 0.122 | revalidation required | -| QM9 | GTE | raw MAE ↓ | 0.071 | revalidation required: official weights | -| OGBG-molhiv | BERT | ROC-AUC ↑ | 82.6% | revalidation required | -| OGBG-molhiv | GTE | ROC-AUC ↑ | 87.4% | revalidation required: official weights | -| Peptides-struct | BERT | Average MAE ↓ | 0.247 | revalidation required | -| Peptides-struct | GTE | Average MAE ↓ | 0.242 | revalidation required: official weights | +| QM9 | BERT | 原始尺度 MAE ↓ | 0.122 | 需要重新验证 | +| QM9 | GTE | 原始尺度 MAE ↓ | 0.071 | 需要使用官方权重重新验证 | +| OGBG-molhiv | BERT | ROC-AUC ↑ | 82.6% | 需要重新验证 | +| OGBG-molhiv | GTE | ROC-AUC ↑ | 87.4% | 需要使用官方权重重新验证 | +| Peptides-struct | BERT | 平均 MAE ↓ | 0.247 | 需要重新验证 | +| Peptides-struct | GTE | 平均 MAE ↓ | 0.242 | 需要使用官方权重重新验证 | -`Paper` is the value reported by the GraphTokenizer paper. GammaGL results must not be compared or described as reproduced until the strict run records the official GTE checkpoint provenance and reports QM9 in raw label units. +“论文结果”是 GraphTokenizer 论文中报告的数值。在严格运行记录官方 GTE 检查点来源并以原始标签单位报告 QM9 结果之前,不得将 GammaGL 结果与论文结果进行比较,也不得称其已经复现。 -Each run keeps result artifacts under `--output-dir`, including the summary JSON, per-run CSV, Markdown/LaTeX paper tables, runtime manifest, checkpoints, and paper-protocol state. JSON remains an output format only; it is not used as an input parameter configuration. +每次运行都会在 `--output-dir` 下保存结果文件,包括汇总 JSON、逐次运行 CSV、Markdown/LaTeX 论文表格、运行时清单、检查点和论文协议状态。JSON 仅作为输出格式,不用于输入参数配置。