CoolFace
Modelpublic

rk-random/PACT-Net

sourceHugging Facemitupdated 9mo agoView on Hugging Face
1likes
train_eval.py278 linesDownload Raw Back to training
1import torch2from sklearn.metrics import mean_squared_error, mean_absolute_error, r2_score3from sklearn.model_selection import KFold4from sklearn.preprocessing import StandardScaler5from tqdm import tqdm6import numpy as np7import os8from datetime import datetime9from pathlib import Path10 11# polyatomic12from torch.amp.autocast_mode import autocast13 14ROOT = Path(__file__).parent.parent.resolve().__str__()15LOG_ROOT = Path(ROOT + "/" + "logs_hyperparameter")16if not os.path.exists(LOG_ROOT):17    os.makedirs(LOG_ROOT, exist_ok=False)18 19 20def setup_log_file(model_name, rep_name, dataset_name):21    from pathlib import Path22    from datetime import datetime23 24    timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")25    fname = f"{model_name}_{rep_name}_{dataset_name}_{timestamp}.txt"26    parent = Path(__file__).parent.parent.resolve().__str__()27    log_dir = Path(parent + "/" + "logs_hyperparameter")28    if not os.path.exists(log_dir):29        os.makedirs(LOG_ROOT, exist_ok=False)30 31    log_path = log_dir / fname32    print(f"[Logging] Writing to: {log_path}")33    return open(log_path, "w")34 35 36def write_log(log_file, text):37    print(text)38    log_file.write(text + "\n")39    log_file.flush()40 41 42def train_gnn_model(model, loader, optimizer, log_file, loss_fn=torch.nn.MSELoss()):43    total_loss = 044    for _ in range(20):45        model.train()46        total_loss = 047        for batch in loader:48            batch = batch.to(next(model.parameters()).device)49            optimizer.zero_grad()50            out = model(batch).squeeze()51            loss = loss_fn(out, batch.y)52            loss.backward()53            optimizer.step()54            total_loss += loss.item()55    avg_loss = total_loss / len(loader)56    write_log(log_file, f"GNN Train Loss: {avg_loss:.4f}")57    return avg_loss58 59 60def eval_gnn_model(model, loader, log_file, scaler, return_preds=False):61    model.eval()62    y_true, y_pred = [], []63    with torch.no_grad():64        for batch in tqdm(loader, desc="Evaluating GNN"):65            batch = batch.to(next(model.parameters()).device)66            out = model(batch).view(-1)67            y_true.append(batch.y.cpu())68            y_pred.append(out.cpu())69    y_true = scaler.inverse_transform(torch.cat(y_true).numpy().reshape(-1, 1)).ravel()70    y_pred = scaler.inverse_transform(torch.cat(y_pred).numpy().reshape(-1, 1)).ravel()71    metrics = report_metrics(y_true, y_pred, log_file)72    if return_preds:73        return metrics, y_true, y_pred74    return metrics75 76 77def train_gp_model(gp_model, X_train, y_train, log_file):78    write_log(log_file, "Training GP model...")79    gp_model.fit(X_train, y_train)80    return gp_model81 82 83def eval_gp_model(gp_model, X_test, y_test, log_file, scaler, return_preds=False):84    write_log(log_file, "Evaluating GP model...")85    y_pred, _ = gp_model.predict(X_test)86    y_test = scaler.inverse_transform(y_test.reshape(-1, 1)).ravel()87    y_pred = scaler.inverse_transform(y_pred.reshape(-1, 1)).ravel()88    metrics = report_metrics(y_test, y_pred, log_file)89    if return_preds:90        return metrics, y_test, y_pred91    return metrics92 93 94def report_metrics(y_true, y_pred, log_file):95    rmse = np.sqrt(mean_squared_error(y_true, y_pred))96    mae = mean_absolute_error(y_true, y_pred)97    r2 = r2_score(y_true, y_pred)98    write_log(log_file, f"RMSE: {rmse:.4f}, MAE: {mae:.4f}, R2: {r2:.4f}")99    return {"rmse": rmse, "mae": mae, "r2": r2}100 101 102def bootstrap_ci(arr, n_boot=1000, ci=95):103    boot_means = [104        np.mean(np.random.choice(arr, size=len(arr), replace=True))105        for _ in range(n_boot)106    ]107    lower = np.percentile(boot_means, (100 - ci) / 2)108    upper = np.percentile(boot_means, 100 - (100 - ci) / 2)109    return np.mean(boot_means), (lower, upper)110 111 112def bootstrap_metric_ci(metric_fn, y_true, y_pred, n_boot=1000, ci=95, rng=None):113    rng = np.random.default_rng(rng)114    y_true = np.asarray(y_true)115    y_pred = np.asarray(y_pred)116    n = len(y_true)117    boot_vals = []118    for _ in range(n_boot):119        idx = rng.choice(n, n, replace=True)120        boot_vals.append(metric_fn(y_true[idx], y_pred[idx]))121    mean_val = np.mean(boot_vals)122    lo, hi = np.percentile(boot_vals, [(100 - ci) / 2, 100 - (100 - ci) / 2])123    return mean_val, (lo, hi)124 125 126def train_polyatomic(127    model, loader, optimizer, loss_fn, scaler_grad, device, scheduler, accum_steps=1128):129    """130    custom training loop for polyatomic GNNs131    uses mixed precision training with autocast132    accum_steps allows gradient accumulation for larger effective batch size133    this was designed for GPU training, but here is in CPU mode134    """135    use_amp = torch.cuda.is_available()136    for _ in range(20):137        model.train()138        total_loss = 0.0139        optimizer.zero_grad()140 141        for i, batch in enumerate(loader):142            batch = batch.to(device)143            batch.x = batch.x.float()144            batch.edge_attr = batch.edge_attr.float()145            batch.graph_feats = batch.graph_feats.float()146            batch.y = batch.y.float()147 148            if use_amp:149                with autocast(device_type="cuda", dtype=torch.float16):150                    output = model(batch)151                    loss = loss_fn(output, batch.y.view(-1)) / accum_steps152            else:153                output = model(batch)154                loss = loss_fn(output, batch.y.view(-1)) / accum_steps155 156            if use_amp:157                scaler_grad.scale(loss).backward()158            else:159                loss.backward()160 161            torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)162 163            if (i + 1) % accum_steps == 0 or (i + 1 == len(loader)):164                if use_amp:165                    scaler_grad.step(optimizer)166                    scaler_grad.update()167                else:168                    optimizer.step()169 170                optimizer.zero_grad()171 172            total_loss += loss.item() * batch.num_graphs * accum_steps173 174        avg_loss = total_loss / len(loader.dataset)175        if scheduler is not None:176            scheduler.step(avg_loss)177 178    return model179 180 181def evaluate_polyatomic(model, loader, device, log_file, scaler, return_preds=False):182    """183    uses autocast for mixed precision evaluation184    this was designed for GPU training, but here is in CPU mode185    """186    model.eval()187    preds, trues = [], []188    with torch.no_grad(), autocast(189        device_type="cpu", dtype=torch.float16190    ):  # change to 'cuda' if using GPU191        for batch in loader:192            batch = batch.to(device)193            out = model(batch)194            preds.append(out.view(-1))195            trues.append(batch.y.view(-1))196    y_pred = torch.cat(preds)197    y_test = torch.cat(trues)198    y_pred = scaler.inverse_transform(y_pred.numpy().reshape(-1, 1)).ravel()199    y_test = scaler.inverse_transform(y_test.numpy().reshape(-1, 1)).ravel()200    metrics = report_metrics(y_test, y_pred, log_file)201    if return_preds:202        return metrics, y_test, y_pred203    return metrics204 205 206def k_fold_eval(207    train_fn,208    eval_fn,209    X_train,210    y_train,211    model_name,212    rep_name,213    dataset_name,214    X_test,215    y_test,216    k=5,217    seed=42,218    log_file=None,219):220    log_file = setup_log_file(model_name, rep_name, dataset_name)221    write_log(log_file, f"Experiment: {model_name}+{rep_name} on {dataset_name}")222 223    kf = KFold(n_splits=k, shuffle=True, random_state=seed)224    fold_metrics, fold_models = [], []225 226    for fold, (tr_idx, val_idx) in enumerate(kf.split(X_train)):227        write_log(log_file, f"\nFOLD {fold+1}/{k}")228        X_tr, X_val = X_train[tr_idx], X_train[val_idx]229        y_tr, y_val = y_train[tr_idx], y_train[val_idx]230 231        scaler = StandardScaler()232        y_tr_s = scaler.fit_transform(y_tr.reshape(-1, 1)).ravel()233        y_val_s = scaler.transform(y_val.reshape(-1, 1)).ravel()234 235        model = train_fn(X_tr, y_tr_s, log_file)236        m = eval_fn(model, X_val, y_val_s, log_file, scaler)237        fold_metrics.append(m)238        fold_models.append((model, m["rmse"]))239 240    write_log(log_file, "\n====== K-FOLD SUMMARY ======")241    for key in fold_metrics[0]:242        vals = [m[key] for m in fold_metrics]243        mean, (lo, hi) = bootstrap_ci(vals)244        write_log(log_file, f"{key.upper()}: {mean:.4f}  [{lo:.4f}, {hi:.4f}]")245 246    best_idx = np.argmin([rm for (_, rm) in fold_models])247    best_model = fold_models[best_idx][0]248    write_log(log_file, f"\n★ Using fold {best_idx+1} model for test inference")249 250    test_scaler = StandardScaler().fit(y_train.reshape(-1, 1))251    y_test_s = test_scaler.transform(y_test.reshape(-1, 1)).ravel()252    test_metrics, y_true_test, y_pred_test = eval_fn(253        best_model, X_test, y_test_s, log_file, test_scaler, return_preds=True254    )255 256    write_log(log_file, f"\n====== HELD-OUT TEST METRICS ======\n{test_metrics}")257    rmse_mean, (rmse_lo, rmse_hi) = bootstrap_metric_ci(258        lambda a, b: np.sqrt(mean_squared_error(a, b)), y_true_test, y_pred_test259    )260    mae_mean, (mae_lo, mae_hi) = bootstrap_metric_ci(261        mean_absolute_error, y_true_test, y_pred_test262    )263    r2_mean, (r2_lo, r2_hi) = bootstrap_metric_ci(r2_score, y_true_test, y_pred_test)264    write_log(265        log_file,266        f"Test RMSE: {rmse_mean:.4f}  (95 % CI: {rmse_lo:.4f}–{rmse_hi:.4f})\n",267    )268    write_log(269        log_file,270        f"Test MAE : {mae_mean :.4f}  (95 % CI: {mae_lo :.4f}–{mae_hi :.4f})\n",271    )272    write_log(273        log_file,274        f"Test R²  : {r2_mean  :.4f}  (95 % CI: {r2_lo  :.4f}–{r2_hi  :.4f})\n",275    )276    log_file.close()277    return fold_metrics, test_metrics278