AI4deeperScience's picture
Add files using upload-large-folder tool
8f4ed7a verified
Raw
History Blame Contribute Delete
28.4 kB
"""Unified training utilities for neural and graph models."""
from __future__ import annotations
import math
import os
from typing import Any, Dict, List, Optional, Tuple, Type
import numpy as np
import torch
import torch.nn as nn
from sklearn.metrics import mean_absolute_error, r2_score
from torch.optim import AdamW
from torch.optim.lr_scheduler import CosineAnnealingLR, LambdaLR, ReduceLROnPlateau
from torch_geometric.loader import DataLoader
from .models import GATModel, HybridModel
from .neural_models import DescriptorNN, FingerprintNN
class NeuralNetworkTrainer:
"""Utility helper to train descriptor or fingerprint networks with scaling."""
def __init__(self, model: nn.Module, device: torch.device, lr: float, weight_decay: float) -> None:
self.model = model.to(device)
self.device = device
self.base_lr = lr
self.optimizer = AdamW(
self.model.parameters(), lr=lr, weight_decay=weight_decay, betas=(0.9, 0.999)
)
self.criterion = nn.SmoothL1Loss(beta=1.0)
self.scheduler: CosineAnnealingLR | None = None
def train_fold(
self,
train_features: np.ndarray,
train_targets: np.ndarray,
val_features: np.ndarray,
val_targets: np.ndarray,
epochs: int,
batch_size: int,
patience: int,
gradient_clip: float,
warmup_epochs: int = 10,
*,
train_lab_indices: Optional[np.ndarray] = None,
val_lab_indices: Optional[np.ndarray] = None,
verbose: bool = True,
) -> Tuple[Dict[str, float], np.ndarray]:
device = self.device
# Target scaling for stability during training
y_mean = float(train_targets.mean())
y_std_raw = float(train_targets.std())
y_std = y_std_raw if y_std_raw > 1e-6 else 1.0
train_y_scaled = (train_targets - y_mean) / y_std
val_y_scaled = (val_targets - y_mean) / y_std
train_features_tensor = torch.tensor(train_features, dtype=torch.float32)
train_targets_tensor = torch.tensor(train_y_scaled, dtype=torch.float32)
val_features_tensor = torch.tensor(val_features, dtype=torch.float32)
val_targets_tensor = torch.tensor(val_y_scaled, dtype=torch.float32)
if train_lab_indices is not None:
train_lab_np = np.asarray(train_lab_indices).reshape(-1)
train_lab_tensor = torch.tensor(train_lab_np, dtype=torch.long)
train_dataset = torch.utils.data.TensorDataset(
train_features_tensor,
train_lab_tensor,
train_targets_tensor,
)
else:
train_dataset = torch.utils.data.TensorDataset(
train_features_tensor,
train_targets_tensor,
)
labs_in_use = train_lab_indices is not None
val_lab_tensor = (
torch.tensor(np.asarray(val_lab_indices).reshape(-1), dtype=torch.long)
if val_lab_indices is not None
else None
)
train_loader = torch.utils.data.DataLoader(
train_dataset, batch_size=batch_size, shuffle=True, drop_last=False
)
if epochs <= warmup_epochs:
warmup_epochs = max(0, epochs - 1)
self.scheduler = CosineAnnealingLR(
self.optimizer,
T_max=max(1, epochs - warmup_epochs),
eta_min=self.base_lr * 0.05,
)
best_state: Dict[str, Any] | None = None
best_mae = float("inf")
patience_counter = 0
for epoch in range(epochs):
self.model.train()
epoch_losses: List[float] = []
for batch in train_loader:
if labs_in_use:
batch_x, batch_lab, batch_y = batch
batch_lab = batch_lab.to(device)
else:
batch_x, batch_y = batch
batch_lab = None
batch_x = batch_x.to(device)
batch_y = batch_y.to(device)
self.optimizer.zero_grad(set_to_none=True)
preds = (
self.model(batch_x, batch_lab)
if batch_lab is not None
else self.model(batch_x)
)
loss = self.criterion(preds, batch_y)
loss.backward()
if gradient_clip:
torch.nn.utils.clip_grad_norm_(self.model.parameters(), gradient_clip)
self.optimizer.step()
epoch_losses.append(loss.item())
# Warmup before cosine decay
if epoch >= warmup_epochs and self.scheduler is not None:
self.scheduler.step()
# Evaluation in original scale
self.model.eval()
with torch.no_grad():
val_preds_scaled = (
self.model(val_features_tensor.to(device), val_lab_tensor.to(device))
if val_lab_tensor is not None
else self.model(val_features_tensor.to(device))
).cpu()
val_preds = val_preds_scaled.numpy() * y_std + y_mean
val_targets_unscaled = val_targets
current_mae = mean_absolute_error(val_targets_unscaled, val_preds)
current_r2 = r2_score(val_targets_unscaled, val_preds)
improved = current_mae + 1e-5 < best_mae
if improved:
best_mae = current_mae
patience_counter = 0
best_state = {
"model": self.model.state_dict(),
"y_mean": y_mean,
"y_std": y_std,
"epoch": epoch + 1,
"mae": current_mae,
"r2": current_r2,
}
else:
patience_counter += 1
if verbose and ((epoch + 1) % 50 == 0 or improved):
lr = self.optimizer.param_groups[0]["lr"]
print(
f" Epoch {epoch+1}: TrainLoss={np.mean(epoch_losses):.4f} "
f"ValMAE={current_mae:.4f} ValR²={current_r2:.4f} LR={lr:.2e}"
f"{' *' if improved else ''}"
)
if patience_counter >= patience:
if verbose:
best_epoch = best_state["epoch"] if best_state else "N/A"
print(f" Early stopping at epoch {epoch+1} (best epoch={best_epoch})")
break
if best_state is None:
raise RuntimeError("Training failed to record a best state.")
# Restore best weights
self.model.load_state_dict(best_state["model"])
self.model.target_mean = best_state["y_mean"]
self.model.target_std = best_state["y_std"]
with torch.no_grad():
final_preds_scaled = (
self.model(val_features_tensor.to(device), val_lab_tensor.to(device))
if val_lab_tensor is not None
else self.model(val_features_tensor.to(device))
).cpu().numpy()
final_preds = final_preds_scaled * self.model.target_std + self.model.target_mean
metrics = {
"r2": best_state["r2"],
"mae": best_state["mae"],
"best_epoch": best_state["epoch"],
}
return metrics, final_preds
def save(self, path: str) -> None:
torch.save(
{
"model_state": self.model.state_dict(),
"target_mean": self.model.target_mean,
"target_std": self.model.target_std,
},
path,
)
def predict(self, features: np.ndarray, lab_indices: Optional[np.ndarray] = None) -> np.ndarray:
self.model.eval()
with torch.no_grad():
feature_tensor = torch.tensor(features, dtype=torch.float32)
if lab_indices is not None:
lab_tensor = torch.tensor(np.asarray(lab_indices).reshape(-1), dtype=torch.long)
preds = self.model(
feature_tensor.to(self.device),
lab_tensor.to(self.device),
).cpu().numpy()
else:
preds = self.model(feature_tensor.to(self.device)).cpu().numpy()
return preds * self.model.target_std + self.model.target_mean
def _ensure_dir(path: str) -> None:
os.makedirs(path, exist_ok=True)
def train_descriptor_nn_fold(
*,
train_features: np.ndarray,
val_features: np.ndarray,
train_targets: np.ndarray,
val_targets: np.ndarray,
fold_idx: int,
device: torch.device,
config: Dict[str, Any],
training_config: Dict[str, Any],
save_dir: str = "oof_models",
train_lab_indices: Optional[np.ndarray] = None,
val_lab_indices: Optional[np.ndarray] = None,
) -> Tuple[DescriptorNN, np.ndarray, Dict[str, float]]:
"""Train Descriptor Neural Network for one CV fold."""
print(f" Training Descriptor NN (Fold {fold_idx})...")
model = DescriptorNN(**config)
trainer = NeuralNetworkTrainer(
model=model,
device=device,
lr=training_config["lr"],
weight_decay=training_config["weight_decay"],
)
metrics, val_predictions = trainer.train_fold(
train_features=train_features,
train_targets=train_targets,
val_features=val_features,
val_targets=val_targets,
epochs=training_config["epochs"],
batch_size=training_config["batch_size"],
patience=training_config["patience"],
gradient_clip=training_config["gradient_clip"],
warmup_epochs=training_config.get("warmup_epochs", 10),
train_lab_indices=train_lab_indices,
val_lab_indices=val_lab_indices,
verbose=True,
)
_ensure_dir(save_dir)
save_path = os.path.join(save_dir, f"desc_nn_fold_{fold_idx}.pt")
trainer.save(save_path)
print(
f" Descriptor NN Fold {fold_idx}: R² = {metrics['r2']:.4f}, "
f"MAE = {metrics['mae']:.4f}, Best Epoch = {metrics['best_epoch']}"
)
return model, val_predictions, metrics
def train_fingerprint_nn_fold(
*,
train_fingerprints: np.ndarray,
val_fingerprints: np.ndarray,
train_targets: np.ndarray,
val_targets: np.ndarray,
fold_idx: int,
device: torch.device,
config: Dict[str, Any],
training_config: Dict[str, Any],
save_dir: str = "oof_models",
train_lab_indices: Optional[np.ndarray] = None,
val_lab_indices: Optional[np.ndarray] = None,
) -> Tuple[FingerprintNN, np.ndarray, Dict[str, float]]:
"""Train Fingerprint Neural Network for one CV fold."""
print(f" Training Fingerprint NN (Fold {fold_idx})...")
model = FingerprintNN(**config)
trainer = NeuralNetworkTrainer(
model=model,
device=device,
lr=training_config["lr"],
weight_decay=training_config["weight_decay"],
)
metrics, val_predictions = trainer.train_fold(
train_features=train_fingerprints,
train_targets=train_targets,
val_features=val_fingerprints,
val_targets=val_targets,
epochs=training_config["epochs"],
batch_size=training_config["batch_size"],
patience=training_config["patience"],
gradient_clip=training_config["gradient_clip"],
warmup_epochs=training_config.get("warmup_epochs", 10),
train_lab_indices=train_lab_indices,
val_lab_indices=val_lab_indices,
verbose=True,
)
_ensure_dir(save_dir)
save_path = os.path.join(save_dir, f"fp_nn_fold_{fold_idx}.pt")
trainer.save(save_path)
print(
f" Fingerprint NN Fold {fold_idx}: R² = {metrics['r2']:.4f}, "
f"MAE = {metrics['mae']:.4f}, Best Epoch = {metrics['best_epoch']}"
)
return model, val_predictions, metrics
def train_gnn_fold(
*,
fold_train_graphs,
fold_val_graphs,
fold_train_lab,
fold_val_lab,
fold_train_targets,
fold_val_targets,
fold_idx: int,
device: torch.device,
config: Dict[str, Any],
training_config: Dict[str, Any],
save_dir: str = "oof_models",
verbose_interval: int = 100,
):
"""Train a single GNN fold with warmup + cosine scheduling."""
_ensure_dir(save_dir)
target_mean = float(np.mean(fold_train_targets))
target_std_raw = float(np.std(fold_train_targets))
target_std = target_std_raw if target_std_raw > 1e-6 else 1.0
scaled_train_targets = (fold_train_targets - target_mean) / target_std
scaled_val_targets = (fold_val_targets - target_mean) / target_std
for i, graph in enumerate(fold_train_graphs):
graph.lab_feature = torch.tensor([fold_train_lab[i]], dtype=torch.long)
graph.y = torch.tensor([scaled_train_targets[i]], dtype=torch.float32)
for i, graph in enumerate(fold_val_graphs):
graph.lab_feature = torch.tensor([fold_val_lab[i]], dtype=torch.long)
graph.y = torch.tensor([scaled_val_targets[i]], dtype=torch.float32)
model = GATModel(**config).to(device)
lr = training_config.get("lr", 3e-4)
weight_decay = training_config.get("weight_decay", 5e-6)
optimizer = AdamW(
model.parameters(),
lr=lr,
weight_decay=weight_decay,
betas=training_config.get("betas", (0.9, 0.999)),
)
criterion = nn.MSELoss()
epochs = training_config.get("epochs", 800)
warmup_epochs = training_config.get("warmup_epochs", 0)
min_lr = training_config.get("min_lr", 1e-6)
factor = training_config.get("factor", 0.7)
def lr_lambda(epoch: int) -> float:
if warmup_epochs > 0 and epoch < warmup_epochs:
return float(epoch + 1) / warmup_epochs
total_decay_epochs = max(1, epochs - warmup_epochs)
progress = max(0.0, epoch - warmup_epochs) / total_decay_epochs
return 0.5 * (1.0 + math.cos(math.pi * progress))
scheduler = LambdaLR(optimizer, lr_lambda)
plateau_scheduler = ReduceLROnPlateau(
optimizer,
patience=training_config.get("plateau_patience", 30),
factor=factor,
min_lr=min_lr,
)
batch_size = training_config.get("batch_size", 32)
gradient_clip = training_config.get("gradient_clip", 1.0)
train_loader = DataLoader(
fold_train_graphs,
batch_size=batch_size,
shuffle=True,
drop_last=True,
)
val_loader = DataLoader(
fold_val_graphs,
batch_size=batch_size,
shuffle=False,
)
best_val_r2 = -float("inf")
patience = training_config.get("patience", 100)
patience_counter = 0
checkpoint_path = os.path.join(save_dir, f"gnn_fold_{fold_idx}.pt")
for epoch in range(epochs):
model.train()
train_loss = 0.0
train_batches = 0
for batch in train_loader:
batch = batch.to(device)
optimizer.zero_grad()
pred = model(
batch.x,
batch.edge_index,
batch.batch,
batch.lab_feature,
getattr(batch, "edge_attr", None),
)
loss = criterion(pred, batch.y)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=gradient_clip)
optimizer.step()
train_loss += loss.item()
train_batches += 1
if train_batches > 0:
train_loss /= train_batches
model.eval()
val_preds_scaled: List[float] = []
val_targets_scaled: List[float] = []
with torch.no_grad():
for batch in val_loader:
batch = batch.to(device)
pred = model(
batch.x,
batch.edge_index,
batch.batch,
batch.lab_feature,
getattr(batch, "edge_attr", None),
)
preds_np = pred.cpu().numpy()
targets_np = batch.y.cpu().numpy()
val_preds_scaled.extend(preds_np.tolist())
val_targets_scaled.extend(targets_np.tolist())
val_preds = np.asarray(val_preds_scaled, dtype=np.float32) * target_std + target_mean
val_targets = np.asarray(val_targets_scaled, dtype=np.float32) * target_std + target_mean
val_r2 = r2_score(val_targets, val_preds)
val_mae = mean_absolute_error(val_targets, val_preds)
scheduler.step()
plateau_scheduler.step(train_loss)
if val_r2 > best_val_r2:
best_val_r2 = val_r2
patience_counter = 0
torch.save(
{
"model_state": model.state_dict(),
"target_mean": target_mean,
"target_std": target_std,
},
checkpoint_path,
)
else:
patience_counter += 1
if verbose_interval and (epoch + 1) % verbose_interval == 0:
current_lr = optimizer.param_groups[0]["lr"]
print(
f" Epoch {epoch + 1}: Train Loss={train_loss:.4f}, "
f"Val R²={val_r2:.4f}, MAE={val_mae:.4f}, LR={current_lr:.2e}"
)
if patience_counter >= patience:
break
if os.path.exists(checkpoint_path):
checkpoint = torch.load(checkpoint_path, map_location=device)
if isinstance(checkpoint, dict) and "model_state" in checkpoint:
model.load_state_dict(checkpoint["model_state"])
target_mean = float(checkpoint.get("target_mean", target_mean))
target_std = float(checkpoint.get("target_std", target_std))
else:
model.load_state_dict(checkpoint)
if target_std == 0:
target_std = 1.0
model.eval()
final_val_preds_scaled: List[float] = []
with torch.no_grad():
for batch in val_loader:
batch = batch.to(device)
pred = model(
batch.x,
batch.edge_index,
batch.batch,
batch.lab_feature,
getattr(batch, "edge_attr", None),
)
final_val_preds_scaled.extend(pred.cpu().numpy().tolist())
final_val_preds = np.asarray(final_val_preds_scaled, dtype=np.float32) * target_std + target_mean
final_r2 = r2_score(fold_val_targets, final_val_preds)
final_mae = mean_absolute_error(fold_val_targets, final_val_preds)
setattr(model, "target_mean", float(target_mean))
setattr(model, "target_std", float(target_std))
print(f" GNN Fold {fold_idx}: R² = {final_r2:.4f}, MAE = {final_mae:.4f}")
return model, final_val_preds.astype(np.float32), {"r2": final_r2, "mae": final_mae}
__all__ = [
"NeuralNetworkTrainer",
"train_descriptor_nn_fold",
"train_fingerprint_nn_fold",
"train_gnn_fold",
"train_hybrid_model_fold",
]
def _prepare_hybrid_graphs(
graphs,
lab_indices: np.ndarray,
targets: np.ndarray,
descriptors: np.ndarray,
) -> None:
for i, graph in enumerate(graphs):
graph.lab_feature = torch.tensor([lab_indices[i]], dtype=torch.long)
graph.y = torch.tensor([targets[i]], dtype=torch.float32)
descriptor_tensor = torch.tensor(descriptors[i], dtype=torch.float32)
if descriptor_tensor.dim() == 1:
descriptor_tensor = descriptor_tensor.unsqueeze(0)
graph.descriptors = descriptor_tensor
def train_hybrid_model_fold(
*,
model_name: str,
graph_model_class: Type[nn.Module],
config: Dict[str, Any],
training_config: Dict[str, Any],
fold_train_graphs,
fold_val_graphs,
fold_train_lab,
fold_val_lab,
fold_train_targets,
fold_val_targets,
fold_train_descriptors,
fold_val_descriptors,
fold_idx: int,
device: torch.device,
save_dir: str = "hybrid_models",
verbose_interval: int = 50,
):
"""Train a Hybrid GNN model (graph + descriptors) for one CV fold."""
_ensure_dir(save_dir)
target_mean = float(np.mean(fold_train_targets))
target_std_raw = float(np.std(fold_train_targets))
target_std = target_std_raw if target_std_raw > 1e-6 else 1.0
scaled_train_targets = (fold_train_targets - target_mean) / target_std
scaled_val_targets = (fold_val_targets - target_mean) / target_std
fold_train_descriptors = np.asarray(fold_train_descriptors, dtype=np.float32)
fold_val_descriptors = np.asarray(fold_val_descriptors, dtype=np.float32)
_prepare_hybrid_graphs(
fold_train_graphs,
fold_train_lab,
scaled_train_targets,
fold_train_descriptors,
)
_prepare_hybrid_graphs(
fold_val_graphs,
fold_val_lab,
scaled_val_targets,
fold_val_descriptors,
)
graph_model_kwargs = dict(config.get("graph_model_kwargs", {}))
graph_feature_dim = config.get("graph_feature_dim")
model = HybridModel(
graph_model_class=graph_model_class,
descriptor_dim=fold_train_descriptors.shape[1],
graph_model_kwargs=graph_model_kwargs,
graph_feature_dim=graph_feature_dim,
descriptor_hidden_dims=config.get("descriptor_hidden_dims"),
final_hidden_dims=config.get("final_hidden_dims"),
dropout=config.get("dropout", 0.2),
use_batch_norm=config.get("use_batch_norm", True),
output_dim=config.get("output_dim", 1),
).to(device)
lr = training_config.get("lr", 3e-4)
weight_decay = training_config.get("weight_decay", 1e-5)
optimizer = AdamW(
model.parameters(),
lr=lr,
weight_decay=weight_decay,
betas=training_config.get("betas", (0.9, 0.999)),
)
criterion = nn.MSELoss()
epochs = training_config.get("epochs", 400)
patience = training_config.get("patience", 80)
gradient_clip = training_config.get("gradient_clip", 1.0)
plateau_scheduler = ReduceLROnPlateau(
optimizer,
patience=training_config.get("plateau_patience", 30),
factor=training_config.get("factor", 0.7),
min_lr=training_config.get("min_lr", 1e-6),
)
batch_size = training_config.get("batch_size", 32)
train_loader = DataLoader(
fold_train_graphs,
batch_size=batch_size,
shuffle=True,
drop_last=len(fold_train_graphs) > batch_size,
)
val_loader = DataLoader(
fold_val_graphs,
batch_size=batch_size,
shuffle=False,
)
best_state: Dict[str, Any] | None = None
best_r2 = -float("inf")
best_mae = float("inf")
patience_counter = 0
for epoch in range(epochs):
model.train()
cumulative_loss = 0.0
batch_count = 0
for batch in train_loader:
batch = batch.to(device)
optimizer.zero_grad()
lab_tensor = batch.lab_feature.squeeze(-1) if batch.lab_feature.dim() > 1 else batch.lab_feature
descriptors_tensor = batch.descriptors.reshape(batch.num_graphs, -1)
preds = model(
batch.x,
batch.edge_index,
batch.batch,
lab_tensor,
descriptors_tensor,
getattr(batch, "edge_attr", None),
)
target_tensor = batch.y.view(-1)
loss = criterion(preds, target_tensor)
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=gradient_clip)
optimizer.step()
cumulative_loss += loss.item()
batch_count += 1
avg_train_loss = cumulative_loss / max(batch_count, 1)
model.eval()
val_preds: List[float] = []
val_targets_list: List[float] = []
val_loss = 0.0
val_batches = 0
with torch.no_grad():
for batch in val_loader:
batch = batch.to(device)
lab_tensor = batch.lab_feature.squeeze(-1) if batch.lab_feature.dim() > 1 else batch.lab_feature
descriptors_tensor = batch.descriptors.reshape(batch.num_graphs, -1)
preds = model(
batch.x,
batch.edge_index,
batch.batch,
lab_tensor,
descriptors_tensor,
getattr(batch, "edge_attr", None),
)
target_tensor = batch.y.view(-1)
val_loss += criterion(preds, target_tensor).item()
val_batches += 1
preds_np = preds.cpu().numpy()
targets_np = target_tensor.cpu().numpy()
val_preds.extend((preds_np * target_std + target_mean).tolist())
val_targets_list.extend((targets_np * target_std + target_mean).tolist())
avg_val_loss = val_loss / max(val_batches, 1)
plateau_scheduler.step(avg_val_loss)
current_r2 = r2_score(val_targets_list, val_preds)
current_mae = mean_absolute_error(val_targets_list, val_preds)
improved = current_r2 > best_r2 + 1e-5
if improved:
best_r2 = current_r2
best_mae = current_mae
patience_counter = 0
best_state = {
"state_dict": model.state_dict(),
"epoch": epoch + 1,
"train_loss": avg_train_loss,
}
torch.save(
{
"model_state": best_state["state_dict"],
"config": config,
"training_config": training_config,
"target_mean": target_mean,
"target_std": target_std,
},
os.path.join(save_dir, f"{model_name}_fold_{fold_idx}.pt"),
)
else:
patience_counter += 1
if verbose_interval and (epoch + 1) % verbose_interval == 0:
current_lr = optimizer.param_groups[0]["lr"]
print(
f" [{model_name}] Epoch {epoch+1}: TrainLoss={avg_train_loss:.4f} "
f"ValLoss={avg_val_loss:.4f} ValR²={current_r2:.4f} ValMAE={current_mae:.4f} LR={current_lr:.2e}"
f"{' *' if improved else ''}"
)
if patience_counter >= patience:
break
if best_state is None:
raise RuntimeError(f"{model_name} fold {fold_idx} failed to improve during training.")
model.load_state_dict(best_state["state_dict"])
model.eval()
final_val_preds: List[float] = []
with torch.no_grad():
for batch in val_loader:
batch = batch.to(device)
lab_tensor = batch.lab_feature.squeeze(-1) if batch.lab_feature.dim() > 1 else batch.lab_feature
descriptors_tensor = batch.descriptors.reshape(batch.num_graphs, -1)
preds = model(
batch.x,
batch.edge_index,
batch.batch,
lab_tensor,
descriptors_tensor,
getattr(batch, "edge_attr", None),
)
preds_np = preds.cpu().numpy()
final_val_preds.extend((preds_np * target_std + target_mean).tolist())
final_val_preds_array = np.asarray(final_val_preds, dtype=np.float32)
metrics = {"r2": best_r2, "mae": best_mae, "best_epoch": best_state["epoch"]}
setattr(model, "target_mean", float(target_mean))
setattr(model, "target_std", float(target_std))
print(
f" {model_name} Fold {fold_idx}: R² = {best_r2:.4f}, MAE = {best_mae:.4f}, "
f"Best Epoch = {best_state['epoch']}"
)
return model, final_val_preds_array, metrics