3. 审核仅训练集预处理与划分
Linux
将随附编写的代码保存为 qm9_gcn_baseline.py。要求 RDKit 和全新数据目录,记录实际数据哈希,固定种子/划分,均值与标准差仅来自训练集目标 0。它是官方接口的组合,不是上游复现结果。
"""Review-only authored baseline using documented PyG APIs; never run in this batch."""
import argparse
import copy
import hashlib
import json
from pathlib import Path
import random
import numpy as np
import rdkit # Require the inspected raw-processing route, not the no-RDKit archive.
import torch
from torch import nn
import torch_geometric
from torch_geometric.datasets import QM9
from torch_geometric.loader import DataLoader
from torch_geometric.nn import GCNConv, global_mean_pool
p = argparse.ArgumentParser()
p.add_argument('data_root')
p.add_argument('output')
a = p.parse_args()
out = Path(a.output)
out.mkdir(parents=True, exist_ok=False)
root = Path(a.data_root)
if root.exists() and any(root.iterdir()):
raise ValueError('Use a fresh data root to prevent an undocumented processed cache')
seed = 37
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.use_deterministic_algorithms(True)
torch.set_num_threads(1)
dataset = QM9(str(root))
# The inspected raw route removes uncharacterized molecules; record actual count.
if len(dataset) < 10000:
raise ValueError('Unexpectedly small QM9 dataset')
ids = torch.randperm(len(dataset), generator=torch.Generator().manual_seed(seed))
# A deliberately small teaching run, not a literature benchmark split.
train_ids, val_ids, test_ids = ids[:8000], ids[8000:9000], ids[9000:10000]
assert not (set(train_ids.tolist()) & set(val_ids.tolist()))
assert not (set(train_ids.tolist()) & set(test_ids.tolist()))
assert not (set(val_ids.tolist()) & set(test_ids.tolist()))
targets = torch.cat([dataset[int(i)].y[:, 0] for i in train_ids])
mu, sigma = targets.mean(), targets.std(unbiased=False)
if not torch.isfinite(targets).all() or not torch.isfinite(sigma) or sigma <= 0:
raise ValueError('Invalid training targets')
loaders = [DataLoader(dataset[ix], batch_size=64, shuffle=(j == 0),
num_workers=0, generator=torch.Generator().manual_seed(seed))
for j, ix in enumerate([train_ids, val_ids, test_ids])]
class Baseline(nn.Module):
def __init__(self):
super().__init__()
self.c1 = GCNConv(dataset.num_node_features, 64)
self.c2 = GCNConv(64, 64)
self.head = nn.Linear(64, 1)
def forward(self, batch):
x = self.c1(batch.x.float(), batch.edge_index).relu()
x = self.c2(x, batch.edge_index).relu()
return self.head(global_mean_pool(x, batch.batch)).flatten()
model = Baseline() # CPU only; no geometry or edge_attr features in this baseline.
optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
@torch.no_grad()
def mae(loader):
model.eval()
total, count = 0.0, 0
for batch in loader:
predicted = model(batch) * sigma + mu
errors = (predicted - batch.y[:, 0]).abs()
if not torch.isfinite(errors).all():
raise ValueError('Nonfinite evaluation')
total += errors.sum().item()
count += batch.num_graphs
return total / count # Debye; no normalization statistics fitted to this split.
history, best, weights = [], float('inf'), None
for epoch in range(1, 11):
model.train()
loss_sum, n = 0.0, 0
for batch in loaders[0]:
optimizer.zero_grad()
loss = nn.functional.mse_loss(model(batch), (batch.y[:, 0] - mu) / sigma)
if not torch.isfinite(loss):
raise ValueError('Nonfinite loss')
loss.backward()
optimizer.step()
loss_sum += loss.item() * batch.num_graphs
n += batch.num_graphs
val = mae(loaders[1])
history.append({'epoch': epoch, 'train_normalized_mse': loss_sum / n, 'val_mae_D': val})
if val < best:
best, weights = val, copy.deepcopy(model.state_dict())
model.load_state_dict(weights)
test_mae = mae(loaders[2]) # Test once after validation-only selection.
torch.save(weights, out / 'best_state_dict.pt')
np.savez(out / 'splits.npz', train=train_ids.numpy(), validation=val_ids.numpy(), test=test_ids.numpy())
files = {str(f.relative_to(root)): hashlib.sha256(f.read_bytes()).hexdigest()
for f in root.rglob('*') if f.is_file()}
report = {'target': 0, 'property': 'dipole moment', 'unit': 'D', 'seed': seed,
'dataset_count': len(dataset), 'split_sizes': [8000, 1000, 1000],
'train_mean_D': mu.item(), 'train_std_D': sigma.item(), 'epochs': 10,
'validation_mae_D': best, 'test_mae_D': test_mae, 'history': history,
'versions': {'torch': torch.__version__, 'pyg': torch_geometric.__version__,
'rdkit': rdkit.__version__, 'numpy': np.__version__},
'device': 'cpu', 'dataset_file_sha256': files,
'scope': 'Teaching baseline; small random split, no bond types or coordinates; no transfer guarantee.'}
(out / 'results.json').write_text(json.dumps(report, indent=2) + '\n')
官方来源