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..6e6426344 100644 --- a/.github/workflows/test_push.yml +++ b/.github/workflows/test_push.yml @@ -1,6 +1,9 @@ name: Build and Test -on: [push, pull_request] +on: + push: + pull_request: + workflow_dispatch: jobs: build-and-test: @@ -22,7 +25,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 +35,14 @@ 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/datasets/test_molecule_benchmark_datasets.py \ + tests/models/test_graph_transformer.py \ + tests/models/test_graph_tokenizer_paper_protocol.py \ + tests/models/test_graph_gte_pretrained.py diff --git a/examples/graph_tokenizer/README.md b/examples/graph_tokenizer/README.md new file mode 100644 index 000000000..752f1a339 --- /dev/null +++ b/examples/graph_tokenizer/README.md @@ -0,0 +1,334 @@ +# GraphTokenizer + +本示例为 GraphTokenizer 提供图序列化、Graph BPE 词元化、掩码语言模型预训练、监督微调、检查点保存和多随机种子评估功能。 + +## 论文 + +GraphTokenizer 由论文[连接图与 Transformer 的图词元化](https://openreview.net/forum?id=jCctxI1BGF)(ICLR 2026)提出。 + +论文实验协议固定到官方实现提交 +`98343a6b025a48fbb6859cd812a12b81ec3ac3cc`。它采用频率引导的欧拉序列化 +(Feuler),每图请求 100 个起点变体(不超过节点数),仅在训练集划分上拟合 +Graph BPE。训练和验证按图随机取一个变体,测试对同一图的全部变体预测取平均; +最佳检查点只由验证集选择,最终报告五次运行的均值和样本标准差(`ddof=1`)。 + +预训练在 BPE 前使用 RandomSwap;微调在 BPE 前使用 RandomSwap 和 +SequenceMasking,并以 0.3 概率向池化特征加入标准差 0.01 的高斯噪声。学习率 +采用线性 warmup 后余弦衰减,最低为初始学习率的 1%。BERT 的位置嵌入表为 +8096,但与官方数据管线相同,实际输入上限固定为 768;GTE 的模型容量为 8192, +论文训练管线的实际输入上限为 8096。 + +Feuler 序列化器在其支持的无向简单图范围内是可逆的。论文中的 GTE +实验会将版本固定的官方编码器加载到原生 TLX GraphGTE 实现中;图词元嵌入和 +任务头均采用全新初始化。 + +Feuler 仅接受无向简单图:不支持自环和平行边,也不接受显式声明 +`directed=True` 或 `is_directed()=True` 的图对象。COO 可以只存储每条无向边 +一次,也可以采用对称形式存储;两种表示都会被规范化为相同的无向边。原始的 +单向 COO 无法表达独立的有向图语义,因此只会被解释为受支持的无向图存储形式。 + +## 数据集 + +以下论文实验命令覆盖三个数据集: + +- QM9:只预测论文默认的 HOMO;训练标签按训练集 z-score 标准化,最终将预测 + 反变换为原始 eV 后计算 MAE。 +- OGBG-molhiv:双 logit、无类别加权的交叉熵,使用 OGB ROC-AUC 报告结果。 +- Peptides-struct:11 目标按训练集逐目标 z-score 标准化,以 L1 训练;评估时 + 反变换到原始标签空间,逐目标计算 MAE 后等权平均。 + +Peptides-struct 的 MLM 阶段使用与官方 `peptides-func` 相同的肽图语料;两套任务 +的图与划分相同,MLM 不读取下游标签,因此本实现直接复用 Peptides-struct 中等价 +的图数据,并在运行清单中记录 `pretraining_graph_corpus=peptides-func`。 + +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 \ + --target-property homo \ + --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 \ + --min-frequency 2 \ + --num-realizations 100 \ + --aggregation-mode avg \ + --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 \ + --target-property homo \ + --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 \ + --min-frequency 2 \ + --num-realizations 100 \ + --aggregation-mode avg \ + --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 \ + --min-frequency 2 \ + --num-realizations 100 \ + --aggregation-mode avg \ + --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 \ + --min-frequency 2 \ + --num-realizations 100 \ + --aggregation-mode avg \ + --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 \ + --min-frequency 2 \ + --num-realizations 100 \ + --aggregation-mode avg \ + --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 \ + --min-frequency 2 \ + --num-realizations 100 \ + --aggregation-mode avg \ + --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..d51fe10b5 --- /dev/null +++ b/examples/graph_tokenizer/graph_tokenizer_trainer.py @@ -0,0 +1,1751 @@ +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 sample_std(values: Sequence[float]) -> float: + values = [float(value) for value in values if not math.isnan(float(value))] + if not values: + return float("nan") + if len(values) == 1: + return 0.0 + avg = sum(values) / len(values) + return math.sqrt( + sum((value - avg) ** 2 for value in values) / (len(values) - 1)) + + +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": sample_std(val_values), + f"test_{metric_name}_mean": mean(test_values), + f"test_{metric_name}_std": sample_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": sample_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": 3, + "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), + "num_realizations": int(getattr(args, "num_realizations", 100)), + } + 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") == 3 + 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) + ], + num_realizations=( + int(getattr(args, "num_realizations", 100)) + if getattr(args, "protocol", None) == "paper" else 1), + ) + 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": 3, + "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() + paper_spec = resolve_dataset_spec(dataset) + protocol.validate_paper_args(args, paper_spec) + protocol.validate_paper_training_options( + protocol.paper_training_options(args, paper_spec), + paper_spec) + 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", {}), + }, + "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", {}), + }, + "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": sample_std(val_values), + f"test_{primary_metric}_mean": mean(values), + f"test_{primary_metric}_std": sample_std(values), + "per_target_mae": aggregate_per_target_mae(runs), + }, + "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="QM9 property; strict paper mode permits only HOMO (the default).", + ) + parser.add_argument("--device", default="cuda", help="Torch device used by --protocol paper.") + parser.add_argument( + "--wait-for-gpu-idle", action="store_true", + help="Wait for the selected CUDA device to have no compute process before paper training.") + 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( + "--num-realizations", type=int, default=100, + help="Paper protocol serializations requested per graph (capped by node count).") + parser.add_argument( + "--aggregation-mode", default="avg", choices=("avg",), + help="Average all serialization-variant predictions per test graph.") + 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..3cf6c2b9c --- /dev/null +++ b/examples/graph_tokenizer/paper_protocol.py @@ -0,0 +1,1979 @@ +"""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 subprocess +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_SOURCE_REVISION = "98343a6b025a48fbb6859cd812a12b81ec3ac3cc" + +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, + "max_position_embeddings": 8192, + "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 not in {None, "homo"}: + raise ValueError( + "QM9 strict paper protocol evaluates the single HOMO target; " + "use --target-property=homo or omit the option.") + 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}.") + if requested_model == "gte" and bool(getattr( + args, "allow_random_gte_init", False)): + raise ValueError( + "Strict paper protocol requires the pinned pretrained GTE weights.") + + +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_token_augmentation() -> Dict[str, Any]: + return { + "pretrain": { + "random_swap": { + "probability": 0.5, "ratio": 0.1, "window_size": 3}}, + "finetune": { + "random_swap": { + "probability": 0.4, "ratio": 0.05, "window_size": 3}, + "sequence_masking": {"probability": 0.3, "ratio": 0.05}, + }, + "evaluation": "none", + } + + +def paper_reference_options(spec, encoder: str) -> Dict[str, Any]: + """Return the pinned official settings used for strict comparison.""" + encoder = str(encoder).lower() + if encoder not in PAPER_ARCHITECTURES: + raise ValueError("Paper protocol encoder must be 'bert' or 'gte'.") + dataset = spec.canonical_name + pretrain_lr = ( + 1e-4 if encoder == "bert" or dataset == "peptides-struct" else 5e-5) + return { + "source_revision": PAPER_SOURCE_REVISION, + "encoder": encoder, + "serialization": "feuler", + "pretrain_epochs": 200, + "finetune_epochs": 200, + "pretrain_lr": pretrain_lr, + "finetune_lr": 5e-5 if dataset == "molhiv" else 1e-5, + "batch_size": 16 if dataset == "peptides-struct" else 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": int( + PAPER_ARCHITECTURES[encoder]["max_position_embeddings"]), + "max_sequence_length": 768 if encoder == "bert" else 8096, + "num_merges": 2000, + "min_frequency": 2, + "num_realizations": 100, + "aggregation_mode": "avg", + "pooling": "mean", + "amp_dtype": "off", + "allow_tf32": False, + "training_loss": "default", + "molhiv_pos_weight": False, + "token_augmentation": _paper_token_augmentation(), + "gaussian_noise": {"probability": 0.3, "std": 0.01}, + "pretraining_graph_corpus": ( + "peptides-func" if dataset == "peptides-struct" else dataset), + } + + +def validate_paper_training_options(options, spec) -> None: + expected = paper_reference_options(spec, options["encoder"]) + for name, expected_value in expected.items(): + actual = options.get(name) + if actual != expected_value: + raise ValueError( + f"Strict paper protocol requires {name}={expected_value!r}; " + f"received {actual!r}.") + + +def paper_training_options(args, spec=None) -> Dict[str, Any]: + encoder = str(args.model).lower() + options = { + "encoder": encoder, + "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, + "max_sequence_length": 768 if encoder == "bert" else 8096, + "num_merges": int(getattr(args, "num_merges", 2000)), + "min_frequency": int(getattr(args, "min_frequency", 2)), + "num_realizations": int(getattr(args, "num_realizations", 100)), + "aggregation_mode": str(getattr(args, "aggregation_mode", "avg")), + "pooling": str(getattr(args, "pooling", "mean")), + "amp_dtype": "off", + "allow_tf32": False, + "training_loss": "default", + "molhiv_pos_weight": False, + } + 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 + }) + if spec is not None: + options.update({ + "source_revision": PAPER_SOURCE_REVISION, + "token_augmentation": _paper_token_augmentation(), + "gaussian_noise": {"probability": 0.3, "std": 0.01}, + "pretraining_graph_corpus": ( + "peptides-func" if spec.canonical_name == "peptides-struct" + else spec.canonical_name), + }) + 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, + max_sequence_length: int | None = None, + num_realizations: int = 1): + max_sequence_length = int( + max_position_embeddings if max_sequence_length is None + else max_sequence_length) + if max_sequence_length <= 0 or max_sequence_length > max_position_embeddings: + raise ValueError( + "max_sequence_length must be positive and no greater than " + "max_position_embeddings.") + num_realizations = int(num_realizations) + if num_realizations <= 0: + raise ValueError("num_realizations must be positive.") + + def truncate(sequence): + sequence = list(map(int, sequence)) + if len(sequence) <= max_sequence_length: + return sequence + return [ + *sequence[:max_sequence_length - 1], + int(tokenizer.special_tokens.sep_token_id), + ] + + def start_nodes(graph): + method = getattr(tokenizer, "realization_start_nodes", None) + if callable(method): + return method(graph, num_realizations) + if num_realizations == 1: + return [None] + value = getattr(graph, "num_nodes", None) + if callable(value): + value = value() + total_nodes = int(value) + actual_samples = min(num_realizations, total_nodes) + step = max(1, total_nodes // actual_samples) + return [(index * step) % total_nodes for index in range(actual_samples)] + + 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.") + results = [] + graph_ids = [] + if num_realizations == 1 and hasattr(tokenizer, "batch_encode_graphs"): + 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] + batch_results = tokenizer.batch_encode_graphs(graph_batch) + if len(batch_results) != len(graph_batch): + raise RuntimeError( + "Tokenizer batch encoding returned the wrong count.") + results.extend(batch_results) + graph_ids.extend(range(start, start + len(graph_batch))) + else: + for graph_id, graph in enumerate(graphs): + for start_node in start_nodes(graph): + result = ( + tokenizer.encode_graph(graph) + if start_node is None + else tokenizer.encode_graph(graph, start_node=start_node)) + results.append(result) + graph_ids.append(graph_id) + sequences = [truncate(result.input_ids) for result in results] + maximum = max((len(sequence) for sequence in sequences), default=1) + global_max_length = max(global_max_length, maximum) + encoded[split_name] = { + "input_ids": sequences, + "serialized_token_ids": [ + list(map(int, getattr(result, "serialized_token_ids", []))) + for result in results + ], + "labels": [list(graphs[graph_id].y) for graph_id in graph_ids], + "graph_ids": graph_ids, + } + return encoded, global_max_length + + +def _load_or_encode_splits( + tokenizer, splits, max_position_embeddings: int, + cache_path, cache_key: str, max_sequence_length: int | None = None, + num_realizations: int = 1): + max_sequence_length = int( + max_position_embeddings if max_sequence_length is None + else max_sequence_length) + 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") == 2 + and cached.get("cache_key") == str(cache_key) + and cached.get("max_position_embeddings") + == int(max_position_embeddings) + and cached.get("max_sequence_length") + == max_sequence_length + and cached.get("num_realizations") + == int(num_realizations)): + 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), + max_sequence_length=max_sequence_length, + num_realizations=num_realizations) + 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": str(cache_key), + "max_position_embeddings": int(max_position_embeddings), + "max_sequence_length": max_sequence_length, + "num_realizations": int(num_realizations), + "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, + max_sequence_length: int | None = None, + num_realizations: int = 1) -> 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": 2, + "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), + "max_sequence_length": int( + max_position_embeddings if max_sequence_length is None + else max_sequence_length), + "num_realizations": int(num_realizations), + } + return hashlib.sha256(pickle.dumps(state, protocol=4)).hexdigest() + + +def _prepare_labels(encoded, spec, target_property): + if spec.canonical_name == "qm9": + source_width = int(spec.output_dim) + source_properties = list(spec.label_keys) + if len(source_properties) != source_width: + raise ValueError( + f"QM9 requires exactly {source_width} target property names; " + f"received {len(source_properties)}.") + selected_target = "homo" if target_property is None else str(target_property) + if selected_target != "homo" or selected_target not in source_properties: + raise ValueError("QM9 paper evaluation requires the HOMO target.") + target_index = source_properties.index(selected_target) + 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) != source_width: + raise ValueError( + f"QM9 {split_name} label row {row_index} must contain " + f"exactly {source_width} targets; received {len(row)}.") + if not math.isfinite(float(row[target_index])): + raise ValueError( + f"QM9 {split_name} label row {row_index} contains a " + "non-finite target.") + train_values = [ + float(row[target_index]) for row in encoded["train"]["labels"]] + means = [sum(train_values) / len(train_values)] + variance = sum( + (value - means[0]) ** 2 for value in train_values + ) / len(train_values) + stds = [max(math.sqrt(variance), 1e-12)] + for split in encoded.values(): + split["labels"] = [ + [(float(row[target_index]) - means[0]) / stds[0]] + for row in split["labels"] + ] + return { + "mean": means, + "std": stds, + "target_properties": [selected_target], + }, 1 + if spec.canonical_name == "peptides-struct": + output_dim = int(spec.output_dim) + target_properties = list(getattr(spec, "label_keys", ())) + if len(target_properties) != output_dim: + target_properties = [f"target_{index}" for index in range(output_dim)] + for split_name, split in encoded.items(): + for row_index, row in enumerate(split["labels"]): + if len(row) != output_dim: + raise ValueError( + f"Peptides-struct {split_name} label row {row_index} " + f"must contain exactly {output_dim} targets; received " + f"{len(row)}.") + train_rows = [list(map(float, row)) for row in encoded["train"]["labels"]] + means, stds = [], [] + for index in range(output_dim): + values = [row[index] for row in train_rows if math.isfinite(row[index])] + if not values: + raise ValueError( + f"Peptides-struct training target {index} has no finite labels.") + target_mean = sum(values) / len(values) + variance = sum( + (value - target_mean) ** 2 for value in values) / len(values) + means.append(target_mean) + stds.append(max(math.sqrt(variance), 1e-12)) + for split in encoded.values(): + split["labels"] = [ + [ + ((float(value) - means[index]) / stds[index] + if math.isfinite(float(value)) else float("nan")) + for index, value in enumerate(row) + ] + for row in split["labels"] + ] + return { + "mean": means, + "std": stds, + "target_properties": target_properties, + }, output_dim + if spec.canonical_name == "molhiv": + return None, 2 + 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 target normalizer.") + return values * stds + means + + +def _metric_semantics(spec, normalizer): + """Describe the training-loss and reported-metric spaces explicitly.""" + if spec.canonical_name in {"qm9", "peptides-struct"}: + 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"}, + } + 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 _random_swap(tokens, probability: float, ratio: float, window_size: int): + if random.random() > float(probability) or len(tokens) <= 1: + return tokens + num_swaps = max(0, int(len(tokens) * float(ratio))) + if num_swaps == 0: + return tokens + augmented = list(tokens) + for _ in range(num_swaps): + center = random.randint(0, len(augmented) - 1) + window_start = max(0, center - int(window_size) // 2) + window_end = min( + len(augmented), center + int(window_size) // 2 + 1) + if window_end - window_start >= 2: + first, second = random.sample(range(window_start, window_end), 2) + augmented[first], augmented[second] = ( + augmented[second], augmented[first]) + return augmented + + +def _sequence_mask(tokens, mask_token_id: int, probability: float, ratio: float): + if random.random() > float(probability) or len(tokens) <= 1: + return tokens + count = max(1, min(int(len(tokens) * float(ratio)), len(tokens) - 1)) + positions = random.sample(range(len(tokens)), count) + augmented = list(tokens) + for position in positions: + augmented[position] = int(mask_token_id) + return augmented + + +def _augment_serialized_tokens(tokens, stage: str, mask_token_id: int): + """Apply the pinned official pre-BPE token transforms.""" + tokens = list(tokens) + if stage == "pretrain": + return _random_swap( + tokens, probability=0.5, ratio=0.1, window_size=3) + if stage == "finetune": + tokens = _random_swap( + tokens, probability=0.4, ratio=0.05, window_size=3) + return _sequence_mask( + tokens, mask_token_id=mask_token_id, + probability=0.3, ratio=0.05) + if stage in {None, "evaluation"}: + return tokens + raise ValueError("augmentation stage must be pretrain, finetune or evaluation.") + + +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, + group_by_graph: bool = False, + return_graph_ids: bool = False, + tokenizer=None, + augmentation_stage: str | None = None, + max_sequence_length: int | None = None): + 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.") + graph_ids = list(encoded_split.get("graph_ids", range(len(input_ids)))) + if len(graph_ids) != len(input_ids): + raise ValueError("Encoded graph ID and input counts must match.") + serialized = list(encoded_split.get( + "serialized_token_ids", [None] * len(input_ids))) + if len(serialized) != len(input_ids): + raise ValueError("Serialized token and input counts must match.") + records = list(zip(input_ids, labels, graph_ids, serialized)) + if group_by_graph: + grouped = {} + for record in records: + grouped.setdefault(int(record[2]), []).append(record) + dataset = list(grouped.values()) + dataset_lengths = [ + max(len(record[0]) for record in group) for group in dataset] + else: + dataset = records + dataset_lengths = [len(sequence) for sequence in input_ids] + + def collate_batch(batch): + if group_by_graph: + batch = [random.choice(group) for group in batch] + prepared = [] + for sequence, label, graph_id, raw_tokens in batch: + if augmentation_stage is not None: + if tokenizer is None or raw_tokens is None: + raise ValueError( + "Token augmentation requires tokenizer and serialized tokens.") + raw_tokens = _augment_serialized_tokens( + raw_tokens, augmentation_stage, + tokenizer.special_tokens.mask_token_id) + sequence = tokenizer.encode_tokens(raw_tokens) + if ( + max_sequence_length is not None + and len(sequence) > int(max_sequence_length)): + sequence = [ + *sequence[:int(max_sequence_length) - 1], + int(tokenizer.special_tokens.sep_token_id), + ] + prepared.append((sequence, label, graph_id)) + batch = prepared + 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) + if return_graph_ids: + return ( + padded, attention_mask, batch_labels, + torch.as_tensor( + [graph_id for _, _, graph_id in batch], dtype=torch.long), + ) + 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, + dataset_lengths, + 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.group_by_graph = bool(group_by_graph) + loader.returns_graph_ids = bool(return_graph_ids) + loader.worker_seed_generator = worker_seed_generator + return loader + + +def _warmup_cosine_scheduler(torch, optimizer, total_steps: int, warmup_ratio: float): + warmup_steps = int(total_steps * float(warmup_ratio)) + if warmup_steps >= total_steps: + warmup_steps = max(0, total_steps - 1) + + def multiplier(step): + if warmup_steps > 0 and step < warmup_steps: + return float(step + 1) / warmup_steps + decay_steps = max(1, total_steps - warmup_steps) + progress = min(1.0, max(0.0, float(step - warmup_steps) / decay_steps)) + cosine = 0.5 * (1.0 + math.cos(math.pi * progress)) + return 0.01 + 0.99 * cosine + + return torch.optim.lr_scheduler.LambdaLR(optimizer, multiplier) + + +def _mask_for_mlm(torch, input_ids, attention_mask, tokenizer, mask_prob: float): + masked = input_ids.clone() + labels = input_ids.clone() + special = torch.zeros_like(labels, dtype=torch.bool) + for token_id in ( + tokenizer.special_tokens.pad_token_id, + tokenizer.special_tokens.cls_token_id, + tokenizer.special_tokens.sep_token_id, + ): + special |= input_ids.eq(int(token_id)) + special |= ~attention_mask.bool() + probabilities = torch.full( + labels.shape, float(mask_prob), dtype=torch.float32, + device=labels.device) + probabilities.masked_fill_(special, 0.0) + selected = torch.bernoulli(probabilities).bool() + if not torch.any(selected): + candidates = (~special).reshape(-1).nonzero( + as_tuple=False).reshape(-1) + if candidates.numel() == 0: + raise ValueError("MLM batch contains no eligible tokens.") + chosen = candidates[torch.randint( + candidates.numel(), (1,), device=labels.device)] + selected.reshape(-1)[chosen] = True + labels[~selected] = -100 + replaced = torch.bernoulli(torch.full( + labels.shape, 0.8, device=labels.device)).bool() & selected + masked[replaced] = int(tokenizer.special_tokens.mask_token_id) + random_replaced = ( + torch.bernoulli(torch.full( + labels.shape, 0.5, device=labels.device)).bool() + & selected & ~replaced) + random_words = torch.randint( + int(tokenizer.max_token_id) + 1, labels.shape, + dtype=torch.long, device=labels.device) + masked[random_replaced] = random_words[random_replaced] + 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 != "default" or pos_weight is not None: + raise ValueError( + "Strict OGBG-molhiv paper loss is unweighted cross entropy.") + if logits.ndim != 2 or logits.shape[-1] != 2: + raise ValueError("OGBG-molhiv paper model must output two logits.") + labels = labels.reshape(-1) + valid = torch.isfinite(labels) + if not torch.any(valid): + raise ValueError("OGBG-molhiv batch contains no valid labels.") + return torch.nn.functional.cross_entropy( + logits[valid], labels[valid].long()) + 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" and spec.canonical_name == "peptides-struct": + return torch.nn.functional.l1_loss(logits, labels) + 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 _visible_cuda_device(torch, device): + logical_index = ( + torch.cuda.current_device() + if device.index is None else int(device.index)) + visible_devices = [value.strip() for value in os.environ.get( + "CUDA_VISIBLE_DEVICES", "").split(",") if value.strip()] + if logical_index < len(visible_devices): + return visible_devices[logical_index] + return str(logical_index) + + +def _wait_for_gpu_idle(torch, device, poll_seconds: int = 60) -> None: + """Wait until the physical CUDA device has no active compute process.""" + if device.type != "cuda": + return + visible_device = _visible_cuda_device(torch, device) + while True: + result = subprocess.run( + ["nvidia-smi", "-i", visible_device, + "--query-compute-apps=pid", "--format=csv,noheader,nounits"], + capture_output=True, text=True, check=True) + active_pids = [line.strip() for line in result.stdout.splitlines() + if line.strip().isdigit()] + if not active_pids: + return + print(json.dumps({ + "event": "paper_wait_for_gpu_idle", + "device": visible_device, + "active_pids": active_pids, + "poll_seconds": int(poll_seconds), + }), flush=True) + time.sleep(int(poll_seconds)) + + +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, target_names=None) -> 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( + target_names if target_names is not None else 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, graph_ids = [], [], [] + non_blocking = bool(getattr(loader, "pin_memory", False)) + with torch.no_grad(): + for batch in loader: + input_ids, attention_mask, labels = batch[:3] + 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.softmax(metric_logits, dim=-1)[:, 1:2].cpu()) + else: + predictions.append(metric_logits.cpu()) + targets.append(labels.cpu()) + if len(batch) > 3: + graph_ids.append(batch[3].cpu()) + y_true = torch.cat(targets, dim=0) + y_pred = torch.cat(predictions, dim=0) + num_sequences = int(len(y_true)) + if graph_ids: + all_graph_ids = torch.cat(graph_ids, dim=0) + unique_ids = [] + seen = set() + for graph_id in all_graph_ids.tolist(): + if graph_id not in seen: + seen.add(graph_id) + unique_ids.append(graph_id) + grouped_targets, grouped_predictions = [], [] + for graph_id in unique_ids: + selected = all_graph_ids.eq(int(graph_id)) + grouped_targets.append(y_true[selected][0]) + grouped_predictions.append(y_pred[selected].mean(dim=0)) + y_true = torch.stack(grouped_targets) + y_pred = torch.stack(grouped_predictions) + average_loss = _require_finite_scalar( + total_loss / max(1, total_examples), "evaluation loss") + target_names = ( + normalizer.get("target_properties") if normalizer is not None else None) + metric_details = compute_paper_metric_details( + spec, y_true.numpy(), y_pred.numpy(), target_names=target_names) + metric_details["metric_space"] = _metric_semantics( + spec, normalizer)["metric_space"] + if normalizer is not None and spec.canonical_name in { + "qm9", "peptides-struct"}: + 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(), + target_names=target_names) + metric_details = { + **raw_details, + "metric_space": _metric_semantics(spec, normalizer)["metric_space"], + } + metric = _require_finite_scalar( + metric_details["metric"], + f"{spec.canonical_name} validation metric", + ) + return { + "loss": average_loss, + "metric": metric, + "num_graphs": int(len(y_true)), + "num_sequences": num_sequences, + **{ + 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": 3, + "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 _forward_supervised( + torch, model, input_ids, attention_mask, gaussian_noise=None): + if gaussian_noise is None or not callable(getattr(model, "_task_logits", None)): + return model(input_ids, attention_mask, task="supervised") + pooled = model(input_ids, attention_mask, task="pooled") + if random.random() <= float(gaussian_noise["probability"]): + pooled = pooled + torch.randn_like(pooled) * float(gaussian_noise["std"]) + return model._task_logits(pooled) + + +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, gaussian_noise=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 = _forward_supervised( + torch, + model, + input_ids.to(device, non_blocking=non_blocking), + attention_mask.to(device, non_blocking=non_blocking), + gaussian_noise=gaussian_noise, + ) + _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): + masked_ids, labels = _mask_for_mlm( + torch, input_ids, attention_mask, tokenizer, mask_prob) + masked_ids = masked_ids.to(device, non_blocking=non_blocking) + labels = labels.to(device, non_blocking=non_blocking) + attention_mask = attention_mask.to( + device, non_blocking=non_blocking) + 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, spec) + validate_paper_training_options(preset, spec) + 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"]), + max_sequence_length=int(preset["max_sequence_length"]), + num_realizations=int(preset["num_realizations"])) + 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, + max_sequence_length=int(preset["max_sequence_length"]), + num_realizations=int(preset["num_realizations"]), + ) + 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")) + if bool(getattr(args, "wait_for_gpu_idle", False)): + _wait_for_gpu_idle(torch, device) + 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", + "tokenizer": tokenizer, + "max_sequence_length": int(preset["max_sequence_length"]), + } + pretrain_loader = _make_loader( + torch, encoded["train"], micro_batch_size, shuffle=True, + bucket_by_length=True, group_by_graph=True, + augmentation_stage="pretrain", **loader_options) + train_loader = _make_loader( + torch, encoded["train"], micro_batch_size, shuffle=True, + bucket_by_length=True, group_by_graph=True, + augmentation_stage="finetune", **loader_options) + val_loader = _make_loader( + torch, encoded["val"], micro_batch_size, shuffle=False, + bucket_by_length=False, group_by_graph=True, **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 = _warmup_cosine_scheduler( + torch, pretrain_optimizer, + _optimizer_steps_per_epoch( + pretrain_loader, accumulation_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, pretrain_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 = _warmup_cosine_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, + gaussian_noise=preset["gaussian_noise"]) + 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, return_graph_ids=True, **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..c8330963d --- /dev/null +++ b/gammagl/models/graph_bert.py @@ -0,0 +1,312 @@ +"""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", "pooled"}: + raise ValueError( + "task must be None, 'mlm', 'supervised', or 'pooled'.") + 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 == "pooled": + return pooled_output + 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..4fa8e1817 --- /dev/null +++ b/gammagl/models/graph_gte.py @@ -0,0 +1,474 @@ +"""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", "pooled"}: + raise ValueError( + "task must be None, 'mlm', 'supervised', or 'pooled'.") + 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 == "pooled": + return pooled_output + 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..50649ad50 --- /dev/null +++ b/gammagl/transforms/graph_serializer.py @@ -0,0 +1,592 @@ +from collections import defaultdict, deque +from dataclasses import dataclass, field +from typing import Any, Dict, Iterable, List, Optional, 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, start_node: Optional[int] = None + ) -> GraphSerializationResult: + graph_view = self._read_graph(graph) + if start_node is not None: + start_node = int(start_node) + if not 0 <= start_node < graph_view["num_nodes"]: + raise ValueError("start_node must identify a node in the graph.") + components = self._connected_components(graph_view["num_nodes"], graph_view["arcs"]) + component_results = [] + + for component in components: + component_start = start_node if start_node in component else None + edges = self._serialize_component( + graph_view, component, start_node=component_start) + tokens = self._edges_to_tokens( + graph_view, edges, component, start_node=component_start) + 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), + "start_node": start_node, + }, + ) + + 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], + start_node: Optional[int] = None) -> 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 = ( + int(start_node) if start_node is not None + else self._select_structural_node( + graph_view, (node for node in component if adjacency.get(node)))) + if not adjacency.get(start): + raise ValueError("start_node has no edge in this component.") + 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], + start_node: Optional[int] = None, + ) -> List[int]: + if not edges: + return self._node_tokens( + graph_view, + int(start_node) if start_node is not None + else 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], + start_node: Optional[int] = None, + ) -> 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 = ( + int(start_node) if start_node is not None + else self._select_structural_node( + graph_view, (node for node in component if adjacency.get(node)))) + if not adjacency.get(start): + raise ValueError("start_node has no edge in this component.") + 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..906e227af --- /dev/null +++ b/gammagl/transforms/graph_tokenizer.py @@ -0,0 +1,411 @@ +from __future__ import annotations + +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.fit_num_realizations = 1 + self.schema_version = 1 + + def fit(self, graphs, graph_ids=None, num_realizations: int = 1): + graphs = list(graphs) + if not graphs: + raise ValueError("Cannot fit GraphTokenizer on an empty training corpus.") + num_realizations = int(num_realizations) + if num_realizations <= 0: + raise ValueError("num_realizations must be positive.") + self.fit_num_realizations = num_realizations + 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 = [] + for graph in graphs: + for start_node in self.realization_start_nodes( + graph, num_realizations): + serialized_sequences.append(self._normalize_serialized_tokens( + self.serializer.serialize( + graph, start_node=start_node).token_ids)) + 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, start_node: int | None = None + ) -> GraphTokenizationResult: + self._require_fitted() + serialized = self.serializer.serialize(graph, start_node=start_node) + 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, + }, + ) + + @staticmethod + def realization_start_nodes(graph: Any, num_realizations: int): + """Match the paper serializer's evenly spaced start-node variants.""" + num_realizations = int(num_realizations) + if num_realizations <= 0: + raise ValueError("num_realizations must be positive.") + if num_realizations == 1: + return [None] + value = ( + graph.get("num_nodes") if isinstance(graph, dict) + else getattr(graph, "num_nodes", None)) + if callable(value): + value = value() + total_nodes = int(value) + if total_nodes <= 0: + return [None] + actual_samples = min(num_realizations, total_nodes) + step = max(1, total_nodes // actual_samples) + return [(index * step) % total_nodes for index in range(actual_samples)] + + 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..6cb286ca2 100644 --- a/tests/data/test_dataset.py +++ b/tests/data/test_dataset.py @@ -5,8 +5,11 @@ import os import pytest +import tensorlayerx as tlx +from gammagl.data import Graph from gammagl.datasets.ppi import PPI +from gammagl.data.dataset import Dataset # dataset record to avoid downloading repeatedly @@ -22,3 +25,21 @@ def test_dataset(): assert len(dataset1) == 20 assert len(dataset2) == 20 + + +def test_torch_dataset_loader_restores_saved_graph(tmp_path): + import torch + + path = tmp_path / 'graph.pt' + graph = Graph( + x=tlx.convert_to_tensor([[1.0], [2.0]]), + edge_index=tlx.convert_to_tensor([[0], [1]]), + ) + torch.save((graph, None), path) + dataset = object.__new__(Dataset) + + restored, slices = dataset.load_data(path) + + assert isinstance(restored, Graph) + assert slices is None + assert tlx.convert_to_numpy(restored.x).tolist() == [[1.0], [2.0]] 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..d5dcaaf83 --- /dev/null +++ b/tests/datasets/test_graph_tokenizer_dataset_download.py @@ -0,0 +1,50 @@ +import importlib.util +import json +import pickle +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_cached_dataset_bundle( + tmp_path, monkeypatch, dataset_name, data_filename): + module = _load_download_module() + bundle_dir = tmp_path / "release" / "data" / dataset_name + bundle_dir.mkdir(parents=True) + (bundle_dir / data_filename).write_bytes(pickle.dumps([dataset_name])) + for split in ("train", "val", "test"): + (bundle_dir / f"{split}_index.json").write_text("[]", encoding="utf-8") + monkeypatch.setenv(module.DATA_BUNDLE_ENV, str(tmp_path / "release")) + + raw_dir = module.materialize_paper_dataset( + dataset_name=dataset_name, + aliases=(), + raw_dir=tmp_path / "datasets" / dataset_name / "raw", + cache_root=tmp_path / "datasets", + allow_download=True, + ) + + assert raw_dir.joinpath(data_filename).is_file() + assert json.loads(raw_dir.joinpath("train_index.json").read_text()) == [] diff --git a/tests/datasets/test_molecule_benchmark_datasets.py b/tests/datasets/test_molecule_benchmark_datasets.py new file mode 100644 index 000000000..acb14eb69 --- /dev/null +++ b/tests/datasets/test_molecule_benchmark_datasets.py @@ -0,0 +1,54 @@ +import gzip +import json +import pickle + +import pytest +import tensorlayerx as tlx + +from gammagl.datasets import OGBGMolHIV, PeptidesStruct, QM9 + + +def _write_raw_dataset(root, name, samples, compressed=False): + raw_dir = root / name / "raw" + raw_dir.mkdir(parents=True) + opener = gzip.open if compressed else open + with opener(raw_dir / ("data.pkl.gz" if compressed else "data.pkl"), "wb") as file: + pickle.dump(samples, file) + for split, indices in {"train": [0], "val": [], "test": []}.items(): + (raw_dir / f"{split}_index.json").write_text(json.dumps(indices)) + + +def test_qm9_loads_a_minimal_local_sample(tmp_path): + _write_raw_dataset(tmp_path, "qm9", [{ + "edge_index": [[0], [1]], "x": [6, 8], "edge_attr": [1], + "properties": {name: 0.0 for name in QM9.label_keys}, + }]) + + dataset = QM9(root=str(tmp_path)) + + assert len(dataset) == 1 + assert tlx.get_tensor_shape(dataset[0].y) == [1, 16] + + +def test_ogbg_molhiv_loads_a_minimal_local_sample(tmp_path): + _write_raw_dataset(tmp_path, "ogbg-molhiv", [( + {"edges": [[0], [1]], "node_type_ids": [6, 8], "edge_type_ids": [1]}, + [1], + )]) + + dataset = OGBGMolHIV(root=str(tmp_path)) + + assert len(dataset) == 1 + assert tlx.get_tensor_shape(dataset[0].y) == [1, 1] + + +def test_peptides_struct_loads_a_minimal_local_sample(tmp_path): + _write_raw_dataset(tmp_path, "peptides-struct", [( + {"edges": [[0], [1]], "node_token_ids": [5, 9], "edge_token_ids": [2]}, + [0.0] * 11, + )], compressed=True) + + dataset = PeptidesStruct(root=str(tmp_path)) + + assert len(dataset) == 1 + assert tlx.get_tensor_shape(dataset[0].y) == [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..6f8e85fb4 --- /dev/null +++ b/tests/models/test_graph_gte_pretrained.py @@ -0,0 +1,56 @@ +import pytest +import tensorlayerx as tlx + +from gammagl.models.graph_gte import GraphGTE +from gammagl.models.graph_gte_pretrained import load_pretrained_encoder + + +torch = pytest.importorskip("torch") +safetensors_torch = pytest.importorskip("safetensors.torch") + + +def _tiny_checkpoint_tensors(model): + layer = model.encoder_layers[0] + return { + "new.embeddings.token_type_embeddings.weight": torch.ones_like( + model.token_type_embeddings.embeddings), + "new.embeddings.LayerNorm.weight": torch.ones_like(model.embedding_norm.gamma), + "new.embeddings.LayerNorm.bias": torch.ones_like(model.embedding_norm.beta), + "new.encoder.layer.0.attention.qkv_proj.weight": torch.ones_like( + layer.attention.qkv_proj.weights), + "new.encoder.layer.0.attention.qkv_proj.bias": torch.ones_like( + layer.attention.qkv_proj.biases), + "new.encoder.layer.0.attention.o_proj.weight": torch.ones_like( + layer.attention.out_proj.weights), + "new.encoder.layer.0.attention.o_proj.bias": torch.ones_like( + layer.attention.out_proj.biases), + "new.encoder.layer.0.attn_ln.weight": torch.ones_like(layer.attention_norm.gamma), + "new.encoder.layer.0.attn_ln.bias": torch.ones_like(layer.attention_norm.beta), + "new.encoder.layer.0.mlp.up_gate_proj.weight": torch.ones_like( + layer.mlp.up_gate_proj.weights), + "new.encoder.layer.0.mlp.down_proj.weight": torch.ones_like( + layer.mlp.down_proj.weights), + "new.encoder.layer.0.mlp.down_proj.bias": torch.ones_like( + layer.mlp.down_proj.biases), + "new.encoder.layer.0.mlp_ln.weight": torch.ones_like(layer.mlp_norm.gamma), + "new.encoder.layer.0.mlp_ln.bias": torch.ones_like(layer.mlp_norm.beta), + } + + +def test_graph_gte_loads_a_local_checkpoint_and_runs_forward(tmp_path): + model = GraphGTE( + vocab_size=8, output_dim=2, hidden_size=16, num_hidden_layers=1, + num_attention_heads=4, intermediate_size=32, max_position_embeddings=16, + ) + checkpoint = tmp_path / "gte.safetensors" + safetensors_torch.save_file(_tiny_checkpoint_tensors(model), str(checkpoint)) + + report = load_pretrained_encoder(model, checkpoint) + logits = model( + tlx.convert_to_tensor([[1, 2, 3]], dtype=tlx.int64), + attention_mask=tlx.convert_to_tensor([[1, 1, 1]], dtype=tlx.int64), + task="supervised", + ) + + assert report["coverage"] == 1.0 + assert tlx.get_tensor_shape(logits) == [1, 2] 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..8b588bf61 --- /dev/null +++ b/tests/models/test_graph_tokenizer_paper_protocol.py @@ -0,0 +1,58 @@ +import importlib.util +from pathlib import Path + +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 + + +def test_paper_protocol_loads_graph_transformer_classes(): + GraphBERT, GraphGTE = _load_paper_protocol_module()._load_paper_model_classes() + + assert GraphBERT.__name__ == "GraphBERT" + assert GraphGTE.__name__ == "GraphGTE" + + +@pytest.mark.parametrize("encoder_type", ["bert", "gte"]) +def test_paper_protocol_creates_a_model_and_runs_forward(encoder_type): + protocol = _load_paper_protocol_module() + model = protocol._create_paper_model( + encoder_type=encoder_type, + vocab_size=16, + pad_token_id=0, + task_type="regression", + output_dim=2, + strict_architecture=False, + allow_random_gte_init=encoder_type == "gte", + model_config={ + "hidden_size": 16, + "num_hidden_layers": 1, + "num_attention_heads": 4, + "intermediate_size": 32, + "max_position_embeddings": 16, + "dropout_rate": 0.0, + }, + ) + input_ids = torch.tensor([[1, 2, 3]], dtype=torch.long) + attention_mask = torch.tensor([[1, 1, 1]], dtype=torch.long) + + supervised = model(input_ids, attention_mask, task="supervised") + mlm = model(input_ids, attention_mask, task="mlm") + + assert supervised.shape == (1, 2) + assert mlm.shape == (1, 3, 16) diff --git a/tests/models/test_graph_transformer.py b/tests/models/test_graph_transformer.py new file mode 100644 index 000000000..569d5d050 --- /dev/null +++ b/tests/models/test_graph_transformer.py @@ -0,0 +1,26 @@ +import pytest +import tensorlayerx as tlx + +from gammagl.models import GraphBERT, GraphGTE + + +@pytest.mark.parametrize("model_class", [GraphBERT, GraphGTE]) +def test_graph_transformer_runs_supervised_and_mlm_forward(model_class): + model = model_class( + vocab_size=16, + 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([[1, 2, 3]], dtype=tlx.int64) + attention_mask = tlx.convert_to_tensor([[1, 1, 1]], dtype=tlx.int64) + + supervised = model(input_ids, attention_mask=attention_mask, task="supervised") + mlm = model(input_ids, attention_mask=attention_mask, task="mlm") + + assert tlx.get_tensor_shape(supervised) == [1, 2] + assert tlx.get_tensor_shape(mlm) == [1, 3, 16] 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..4cc20b0af --- /dev/null +++ b/tests/transforms/test_graph_bpe_cpp_install.py @@ -0,0 +1,10 @@ +import pytest + + +def test_graph_bpe_cpp_extension_trains_a_token_sequence(): + native = pytest.importorskip("third_party.graph_bpe_cpp._graph_bpe") + + result = native.train_bpe([[1, 2, 1, 2]], 1, 2) + + assert result["merge_rules"] + assert result["vocab_size"] >= 3 diff --git a/tests/transforms/test_graph_tokenizer.py b/tests/transforms/test_graph_tokenizer.py new file mode 100644 index 000000000..9c456f876 --- /dev/null +++ b/tests/transforms/test_graph_tokenizer.py @@ -0,0 +1,43 @@ +from gammagl.transforms import ( + FrequencyGuidedEulerianSerializer, + GraphBPE, + GraphTokenizer, +) + + +GRAPH = { + "edge_index": [[0, 1], [1, 2]], + "x": [6, 8, 7], + "edge_attr": [1, 2], + "num_nodes": 3, +} + + +def test_feuler_serializer_serializes_a_graph(): + serializer = FrequencyGuidedEulerianSerializer().fit([GRAPH]) + + result = serializer.serialize(GRAPH) + + assert result.token_ids + assert result.metadata["num_nodes"] == 3 + + +def test_graph_bpe_encodes_and_decodes_token_sequence(): + bpe = GraphBPE(num_merges=1, min_frequency=2).fit([[1, 2, 1, 2]]) + + encoded = bpe.encode([1, 2, 1, 2]) + + assert bpe.codebook.merge_rules + assert bpe.decode(encoded) == [1, 2, 1, 2] + + +def test_graph_tokenizer_fits_and_encodes_a_graph(): + tokenizer = GraphTokenizer( + bpe=GraphBPE(num_merges=2, min_frequency=1)).fit([GRAPH]) + + encoding = tokenizer.encode_graph(GRAPH) + + assert encoding.input_ids[0] == tokenizer.special_tokens.cls_token_id + assert encoding.input_ids[-1] == tokenizer.special_tokens.sep_token_id + assert len(encoding.input_ids) == len(encoding.attention_mask) + assert encoding.serialized_token_ids 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"], + ) + ], +)