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..18dd59596 --- /dev/null +++ b/examples/graph_tokenizer/README.md @@ -0,0 +1,298 @@ +# GraphTokenizer + +本示例为 GraphTokenizer 提供图序列化、Graph BPE 词元化、掩码语言模型预训练、监督微调、检查点保存和多随机种子评估功能。 + +## 论文 + +GraphTokenizer 由论文[连接图与 Transformer 的图词元化](https://openreview.net/forum?id=jCctxI1BGF)(ICLR 2026)提出。 + +论文实验协议采用频率引导的欧拉序列化(Feuler),仅在训练集划分上拟合 Graph BPE,使用掩码语言模型进行预训练,根据验证集性能选择最佳微调检查点,只在测试集上评估一次,并报告五次运行结果的均值和总体标准差。 + +Feuler 序列化器在其支持的无向简单图范围内是可逆的。论文中的 GTE +实验会将版本固定的官方编码器加载到原生 TLX GraphGTE 实现中;图词元嵌入和 +任务头均采用全新初始化。 + +Feuler 仅接受无向简单图:不支持自环和平行边,也不接受显式声明 +`directed=True` 或 `is_directed()=True` 的图对象。COO 可以只存储每条无向边 +一次,也可以采用对称形式存储;两种表示都会被规范化为相同的无向边。原始的 +单向 COO 无法表达独立的有向图语义,因此只会被解释为受支持的无向图存储形式。 + +## 数据集 + +以下论文实验命令覆盖三个数据集: + +- QM9:联合 16 目标回归,使用 MAE 报告结果。 +- OGBG-molhiv:二分类,使用 OGB ROC-AUC 报告结果。 +- Peptides-struct:11 目标回归,使用平均 MAE 报告结果。 + +GammaGL 数据集类会在首次使用时下载 GraphTokenizer 官方发布的数据包。下载的 +归档文件会在解压前使用官方公布的 SHA-256 进行校验: + +``` +5c437c3c0d4b7278379c0e70d57f98148e5c815d753d8cf68e2a45952bcce459 +``` + +数据包会缓存在 `/.graph_tokenizer_release` 下,然后复制到各数据集的 +常规 `raw/` 目录。通过 +`GAMMAGL_GRAPH_TOKENIZER_DATA_BUNDLE=/path/to/bundle` 指定的本地归档文件也会 +使用相同的摘要进行校验。将该变量设置为目录表示显式启用本地开发覆盖;程序 +不会从用户提供的 URL 下载文件。只有自动下载的远程数据包通过校验后,程序才会 +反序列化其中发布的 `data.pkl(.gz)` 文件。 + +## 环境要求 + +论文模式要求: + +- `TL_BACKEND=torch` +- PyTorch 2.1.2,配套 CUDA 12.1 +- DGL 2.4.0,配套 CUDA 12.1 +- PyTorch Geometric 2.4.0 +- TensorLayerX、NumPy、OGB 和 GammaGL 软件包 +- 用于固定版本原生 TLX GTE 检查点转换的 `huggingface-hub` 和 `safetensors` +- 执行以下命令所需的原生 Graph BPE 扩展 + +使用以下方式安装依赖,不会将论文专用软件包加入 GammaGL 的核心依赖集合: + +```bash +pip install -e '.[graph-tokenizer-paper]' +# 也可以在版本固定的论文环境中使用 examples/graph_tokenizer/requirements.txt。 +``` + +`transformers` 是可选依赖:它仅用于检查点和参考实现等价性测试,不会用于正式的 +TLX GraphBERT/GraphGTE 训练前向传播路径。 + +在仓库根目录编译原生 Graph BPE 后端: + +```bash +export TL_BACKEND=torch +export DGLBACKEND=pytorch +python third_party/graph_bpe_cpp/setup.py build_ext --inplace +``` + +论文模型直接使用 TensorLayerX 的公共层和运算;GraphBERT 和 GraphGTE 不依赖 Hugging Face Transformers。 + +## 运行方法 + +所有论文超参数均通过 `argparse` 显式传入。以下命令使用相对于仓库的数据、缓存和结果路径,可以直接在 GammaGL 仓库根目录运行。 + +### 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 +``` + +在开始长时间运行前,可以向相应命令添加 `--preflight`,在不训练的情况下检查数据集、模型、运行环境和 BPE 后端。`--resume` 会从 `last_state.pt` 恢复最近的训练阶段、优化器、调度器、随机数生成器和 AMP 缩放器状态;`best.pt` 仍然是轻量级的最佳模型检查点。 + +执行一个非论文模式的小型冒烟测试: + +```bash +python examples/graph_tokenizer/graph_tokenizer_trainer.py \ + --smoke \ + --dataset qm9 \ + --model bert \ + --data-root data \ + --bpe-backend python +``` + +## 实验结果 + +论文报告的五次运行平均结果如下: + +| 数据集 | 编码器 | 指标 | 论文结果 | GammaGL 状态 | +| --- | --- | --- | ---: | --- | +| 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 | 需要使用官方权重重新验证 | + +“论文结果”是 GraphTokenizer 论文中报告的数值。在严格运行记录官方 GTE 检查点来源并以原始标签单位报告 QM9 结果之前,不得将 GammaGL 结果与论文结果进行比较,也不得称其已经复现。 + +每次运行都会在 `--output-dir` 下保存结果文件,包括汇总 JSON、逐次运行 CSV、Markdown/LaTeX 论文表格、运行时清单、检查点和论文协议状态。JSON 仅作为输出格式,不用于输入参数配置。 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"], + ) + ], +)