OSCR

A benchmark dataset and interpretable deep learning framework for drug-induced developmental neurotoxicity prediction.

Code ↔ Paper

13 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 13 matches · 1 of them tie a paragraph to a whole file, not to given lines: a weak match, whose lines are not tinted
  1. [1] § Materials and methods › Model construction and training strategy › Model training and evaluation ↔ 3-DNN_MACCS_modeling/DNN-test_selection.py, lines 33–67 · score 0.70 · monitoring validation, model selection, cross validation, weights, optimal, epochs
  2. [2] § Materials and methods › Performance metrics and statistical analysis ↔ 3-DNN_MACCS_modeling/DNN-independent validation.py, lines 208–293 · score 0.69 · F1 score, Matthew, TN, TP, FN, precision
  3. [3] § Materials and methods › Performance metrics and statistical analysis ↔ 3-DNN_MACCS_modeling/DNN-test_selection.py, lines 560–599 · score 0.67 · predicted toxic, evaluation metrics, F1 score, precision, recall, accuracy
  4. [4] § Materials and methods › Model construction and training strategy › Model training and evaluation ↔ 3-DNN_MACCS_modeling/DNN-training-5-fold-CV.py, lines 34–65 · score 0.65 · fold cross validation, network architectures, overfitting, epochs, training, Model
  5. [5] § Results and discussion › Analysis of the training stability and generalization ability of the DNN_MACCS model ↔ 3-DNN_MACCS_modeling/DNN-test_selection.py, lines 889–968 · score 0.64 · developmental neurotoxicity prediction, validation AUC, F1 score, dnn maccs, overfitting, configuration
  6. [6] § Materials and methods › Model interpretability analysis ↔ 4-applicability_domain_SHAP_analysis/SHAP-analysisl.py, lines 1056–1121 · score 0.62 · SHapley, exPlanations, Model interpretability, Additive, global, transparency
  7. [7] § Materials and methods › Benchmark dataset preparation › Data partitioning of the benchmark dataset ↔ 1-Benchmark_Dataset_Preparation/plot_umap_chemical_space.py, the whole file · a weak match · score 0.59 · chemical space, DNT positive, UMAP, benchmark, SMILES, fingerprints
  8. [8] § Results and discussion › Strengths and limitations ↔ 4-applicability_domain_SHAP_analysis/Application_Domain_UMAP_tSNE.py, lines 429–462 · score 0.58 · broader chemical, applicability domain, experimentally validated, model generalizability, chemical space, expand
  9. [9] § Results and discussion › Analysis of the training stability and generalization ability of the DNN_MACCS model ↔ 3-DNN_MACCS_modeling/DNN-test_selection.py, lines 889–968 · score 0.56 · potential overfitting, best validation, dnn maccs, loss, epochs, accuracy
  10. [10] § Materials and methods › Benchmark dataset preparation › Preprocessing of positive samples ↔ 1-Benchmark_Dataset_Preparation/DNTREF_curation.py, lines 25–46 · score 0.54 · heavy atoms, mixtures, parent, curation, SMILES
  11. [11] § Materials and methods › Chemical space visualization and applicability domain analysis ↔ 4-applicability_domain_SHAP_analysis/Application_Domain_UMAP_tSNE.py, lines 198–201 · score 0.52 · Dimensionality reduction, UMAP, Neighbor, SNE, domain
  12. [12] § Materials and methods › Benchmark dataset preparation ↔ 3-DNN_MACCS_modeling/DNN-test_selection.py, lines 1014–1134 · score 0.52 · model training, developmental neurotoxicity prediction, pipeline, selection, preprocessing
  13. [13] § Materials and methods › Benchmark dataset preparation › Preprocessing of presumed negative samples and construction of a balanced benchmark dataset ↔ 1-Benchmark_Dataset_Preparation/hard_negative_matching.py, lines 54–90 · score 0.50 · hard negative matching, greedy, MW, log, benchmark

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Python · 1,147 lines · 51 KB · no license · 5 matches

  1. """
  2. Developmental Neurotoxicity Prediction DNN QSAR Model
  3. Author: DeepSeek Assistant
  4. Function: Train DNN model using validation accuracy for model selection
  5. Input files: DNT_benchmark_random_train.csv, DNT_benchmark_random_val.csv
  6. Output: Trained model, performance evaluation results, visualization charts
  7. """
  8. import numpy as np
  9. import pandas as pd
  10. import matplotlib
  11. matplotlib.use('TkAgg') # Set backend to avoid GUI issues
  12. import matplotlib.pyplot as plt
  13. import seaborn as sns
  14. from sklearn.metrics import (roc_auc_score, accuracy_score, precision_score,
  15. recall_score, f1_score, confusion_matrix,
  16. roc_curve, classification_report)
  17. from sklearn.preprocessing import StandardScaler
  18. from rdkit import Chem
  19. from rdkit.Chem import MACCSkeys
  20. import tensorflow as tf
  21. from tensorflow import keras
  22. from tensorflow.keras import layers, models, optimizers, regularizers, callbacks
  23. import warnings
  24. warnings.filterwarnings('ignore')
  25. import os
  26. # Set random seeds for reproducibility
  27. SEED = 42
  28. np.random.seed(SEED)
  29. tf.random.set_seed(SEED)
  30. # ==================== Configuration ====================
  31. class Config:
  32. """Model configuration parameters"""
  33. # Optimal parameters (based on previous cross-validation results)
  34. OPTIMAL_PARAMS = {
  35. 'layer_sizes': [256, 128, 64],
  36. 'dropout_rate': 0.4,
  37. 'learning_rate': 0.0005,
  38. 'l2_reg': 0.001,
  39. 'batch_size': 32,
  40. 'epochs': 200 # Slightly reduced for practical training time
  41. }
  42. # Training parameters
  43. REDUCE_LR_PATIENCE = 20
  44. REDUCE_LR_FACTOR = 0.5
  45. MIN_LR = 1e-6
  46. # Early stopping parameters
  47. EARLY_STOPPING_MONITOR = 'val_accuracy' # Monitor validation accuracy
  48. EARLY_STOPPING_PATIENCE = 40 # Stop if no improvement for 40 epochs
  49. EARLY_STOPPING_MIN_DELTA = 0.001 # Minimum change to qualify as improvement
  50. EARLY_STOPPING_MODE = 'max' # Maximize validation accuracy
  51. EARLY_STOPPING_RESTORE_BEST_WEIGHTS = True # Restore best weights when stopped
  52. # Model selection
  53. MODEL_SELECTION_METRIC = 'val_accuracy' # Use validation accuracy for model selection
  54. MODEL_SELECTION_MODE = 'max' # Maximize validation accuracy
  55. # Visualization settings
  56. FIGURE_DPI = 300
  57. # File paths
  58. TRAIN_FILE = 'DNT_benchmark_random_train-4_split_train.csv'
  59. VAL_FILE = 'DNT_benchmark_random_train-4_split_test.csv'
  60. # ==================== Data Preprocessing ====================
  61. class DataProcessor:
  62. """Data preprocessing class"""
  63. @staticmethod
  64. def load_data(train_path, val_path):
  65. """Load training and validation datasets"""
  66. try:
  67. train_df = pd.read_csv(train_path)
  68. val_df = pd.read_csv(val_path)
  69. print("=" * 60)
  70. print("Data Loading Complete")
  71. print("=" * 60)
  72. print(f"Training set size: {len(train_df)}")
  73. print(f"Validation set size: {len(val_df)}")
  74. print(f"\nTraining set class distribution:")
  75. print(train_df['label'].value_counts())
  76. print(f"\nValidation set class distribution:")
  77. print(val_df['label'].value_counts())
  78. if 'source' in train_df.columns:
  79. print(f"\nTraining set source distribution:")
  80. print(train_df['source'].value_counts())
  81. return train_df, val_df
  82. except Exception as e:
  83. print(f"Failed to load data: {e}")
  84. raise
  85. @staticmethod
  86. def smiles_to_maccs(smiles):
  87. """Convert SMILES to MACCS fingerprints"""
  88. try:
  89. mol = Chem.MolFromSmiles(smiles)
  90. if mol is not None:
  91. fp = MACCSkeys.GenMACCSKeys(mol)
  92. return np.array(fp)
  93. else:
  94. return np.zeros(167)
  95. except:
  96. return np.zeros(167)
  97. def prepare_features(self, train_df, val_df):
  98. """Prepare feature matrices and label vectors"""
  99. print("\nGenerating MACCS fingerprint features...")
  100. # Generate MACCS fingerprints
  101. X_train = np.array([self.smiles_to_maccs(s) for s in train_df['smiles']])
  102. y_train = train_df['label'].values
  103. X_val = np.array([self.smiles_to_maccs(s) for s in val_df['smiles']])
  104. y_val = val_df['label'].values
  105. # Check feature dimensions
  106. print(f"Training set feature dimension: {X_train.shape}")
  107. print(f"Validation set feature dimension: {X_val.shape}")
  108. # Data standardization
  109. scaler = StandardScaler()
  110. X_train_scaled = scaler.fit_transform(X_train)
  111. X_val_scaled = scaler.transform(X_val)
  112. print("Feature standardization complete")
  113. return X_train_scaled, y_train, X_val_scaled, y_val, scaler
  114. # ==================== DNN Model Builder ====================
  115. class DNNModelBuilder:
  116. """DNN model builder"""
  117. @staticmethod
  118. def create_model(input_dim=167,
  119. layer_sizes=[256, 128, 64],
  120. dropout_rate=0.3,
  121. learning_rate=0.001,
  122. l2_reg=0.001):
  123. """Create DNN model architecture"""
  124. model = models.Sequential()
  125. # Input layer and first hidden layer
  126. model.add(layers.Dense(layer_sizes[0], input_dim=input_dim,
  127. activation='relu',
  128. kernel_regularizer=regularizers.l2(l2_reg),
  129. name=f"dense_input_{layer_sizes[0]}"))
  130. model.add(layers.BatchNormalization(name=f"bn_1"))
  131. model.add(layers.Dropout(dropout_rate, name=f"dropout_1"))
  132. # Add subsequent hidden layers
  133. for i, layer_size in enumerate(layer_sizes[1:], start=2):
  134. model.add(layers.Dense(layer_size, activation='relu',
  135. kernel_regularizer=regularizers.l2(l2_reg),
  136. name=f"dense_hidden_{i}_{layer_size}"))
  137. model.add(layers.BatchNormalization(name=f"bn_{i}"))
  138. model.add(layers.Dropout(dropout_rate, name=f"dropout_{i}"))
  139. # Output layer (binary classification)
  140. model.add(layers.Dense(1, activation='sigmoid', name="output"))
  141. # Compile model
  142. optimizer = optimizers.Adam(learning_rate=learning_rate)
  143. model.compile(loss='binary_crossentropy',
  144. optimizer=optimizer,
  145. metrics=['accuracy',
  146. keras.metrics.AUC(name='auc'),
  147. keras.metrics.Precision(name='precision'),
  148. keras.metrics.Recall(name='recall')])
  149. return model
  150. # ==================== Model Trainer ====================
  151. class ModelTrainer:
  152. """Model trainer with validation accuracy selection and early stopping"""
  153. def __init__(self, config):
  154. self.config = config
  155. self.best_epoch = 0
  156. self.best_val_accuracy = 0.0
  157. self.best_model_path = None
  158. self.early_stopped_epoch = None
  159. self.stopped_reason = None
  160. def train_final_model(self, X_train, y_train, X_val, y_val):
  161. """Train final model using optimal parameters with validation accuracy selection and early stopping"""
  162. print("\n" + "=" * 60)
  163. print("Training Final Model with Optimal Parameters")
  164. print("=" * 60)
  165. print("Model Selection Strategy: Using validation accuracy to select best model")
  166. print("Early Stopping: Enabled with patience =", self.config.EARLY_STOPPING_PATIENCE)
  167. print("=" * 60)
  168. # Get optimal parameters
  169. optimal_params = self.config.OPTIMAL_PARAMS
  170. # Separate model building parameters and training parameters
  171. model_build_params = {
  172. 'layer_sizes': optimal_params['layer_sizes'],
  173. 'dropout_rate': optimal_params['dropout_rate'],
  174. 'learning_rate': optimal_params['learning_rate'],
  175. 'l2_reg': optimal_params['l2_reg']
  176. }
  177. # Create final model
  178. final_model = DNNModelBuilder.create_model(
  179. input_dim=X_train.shape[1],
  180. **model_build_params
  181. )
  182. # Print model architecture
  183. print("\nModel Architecture:")
  184. final_model.summary()
  185. # Create directory for saved models
  186. os.makedirs('checkpoints', exist_ok=True)
  187. self.best_model_path = 'checkpoints/best_model_val_acc.h5'
  188. # Training callbacks - focus on validation accuracy with early stopping
  189. callbacks_list = [
  190. # Early stopping callback
  191. callbacks.EarlyStopping(
  192. monitor=self.config.EARLY_STOPPING_MONITOR,
  193. patience=self.config.EARLY_STOPPING_PATIENCE,
  194. min_delta=self.config.EARLY_STOPPING_MIN_DELTA,
  195. mode=self.config.EARLY_STOPPING_MODE,
  196. restore_best_weights=self.config.EARLY_STOPPING_RESTORE_BEST_WEIGHTS,
  197. verbose=1
  198. ),
  199. # Learning rate reduction
  200. callbacks.ReduceLROnPlateau(
  201. monitor='val_loss',
  202. factor=self.config.REDUCE_LR_FACTOR,
  203. patience=self.config.REDUCE_LR_PATIENCE,
  204. min_lr=self.config.MIN_LR,
  205. verbose=1
  206. ),
  207. # Model checkpoint - save best model based on validation accuracy
  208. callbacks.ModelCheckpoint(
  209. self.best_model_path,
  210. monitor=self.config.MODEL_SELECTION_METRIC,
  211. mode=self.config.MODEL_SELECTION_MODE,
  212. save_best_only=True,
  213. save_weights_only=False,
  214. verbose=1
  215. ),
  216. callbacks.CSVLogger('training_log.csv'),
  217. # Custom callback to track best epoch
  218. callbacks.LambdaCallback(
  219. on_epoch_end=lambda epoch, logs: self._update_best_epoch(epoch, logs)
  220. ),
  221. # Custom callback to track early stopping
  222. callbacks.LambdaCallback(
  223. on_epoch_end=lambda epoch, logs: self._check_early_stopping(epoch, logs)
  224. )
  225. ]
  226. # Train model
  227. print(f"\nTraining Configuration:")
  228. print(f" Network Architecture: {model_build_params['layer_sizes']}")
  229. print(f" Dropout Rate: {model_build_params['dropout_rate']}")
  230. print(f" Learning Rate: {model_build_params['learning_rate']}")
  231. print(f" L2 Regularization: {model_build_params['l2_reg']}")
  232. print(f" Batch Size: {optimal_params['batch_size']}")
  233. print(f" Maximum Epochs: {optimal_params['epochs']}")
  234. print(f" Early Stopping Patience: {self.config.EARLY_STOPPING_PATIENCE}")
  235. print(f" Model Selection Metric: {self.config.MODEL_SELECTION_METRIC}")
  236. print(f" Best Model Saved To: {self.best_model_path}")
  237. history = final_model.fit(
  238. X_train, y_train,
  239. validation_data=(X_val, y_val),
  240. epochs=optimal_params['epochs'],
  241. batch_size=optimal_params['batch_size'],
  242. callbacks=callbacks_list,
  243. verbose=1
  244. )
  245. # Check if early stopping occurred
  246. if hasattr(final_model, 'history') and len(final_model.history.history['loss']) < optimal_params['epochs']:
  247. self.early_stopped_epoch = len(final_model.history.history['loss'])
  248. self.stopped_reason = f"Early stopping triggered at epoch {self.early_stopped_epoch}"
  249. print(f"\n⚠ Early stopping was triggered at epoch {self.early_stopped_epoch}")
  250. print(f" Best validation accuracy: {self.best_val_accuracy:.4f}")
  251. # Load the best model based on validation accuracy
  252. print(f"\nLoading best model from epoch {self.best_epoch + 1}")
  253. print(f"Best validation accuracy: {self.best_val_accuracy:.4f}")
  254. # Create a new model instance and load the best weights
  255. best_model = DNNModelBuilder.create_model(
  256. input_dim=X_train.shape[1],
  257. **model_build_params
  258. )
  259. best_model.load_weights(self.best_model_path)
  260. return best_model, history, final_model
  261. def _update_best_epoch(self, epoch, logs):
  262. """Update the best epoch based on validation accuracy"""
  263. if logs is not None and 'val_accuracy' in logs:
  264. current_val_acc = logs['val_accuracy']
  265. if current_val_acc > self.best_val_accuracy:
  266. self.best_val_accuracy = current_val_acc
  267. self.best_epoch = epoch
  268. print(f"\nNew best validation accuracy: {current_val_acc:.4f} at epoch {epoch + 1}")
  269. def _check_early_stopping(self, epoch, logs):
  270. """Check if early stopping criteria are met"""
  271. # This is handled by Keras EarlyStopping callback
  272. pass
  273. @staticmethod
  274. def plot_training_history(history, params, best_epoch=None, best_val_acc=None):
  275. """Plot training history with best epoch marked and y-axis starting from 0"""
  276. plt.figure(figsize=(14, 5))
  277. # Subplot 1: Loss (y-axis starts from 0)
  278. plt.subplot(1, 3, 1)
  279. plt.plot(history.history['loss'], label='Training Loss')
  280. plt.plot(history.history['val_loss'], label='Validation Loss')
  281. # Mark best epoch if provided
  282. if best_epoch is not None:
  283. plt.axvline(x=best_epoch, color='r', linestyle='--', alpha=0.7,
  284. label=f'Best Epoch: {best_epoch + 1}')
  285. plt.plot(best_epoch, history.history['val_loss'][best_epoch],
  286. 'ro', markersize=8)
  287. plt.title('Model Loss')
  288. plt.xlabel('Epochs')
  289. plt.ylabel('Loss')
  290. plt.legend()
  291. plt.grid(True, alpha=0.3)
  292. plt.ylim(bottom=0) # Set y-axis to start from 0
  293. # Subplot 2: AUC (y-axis starts from 0)
  294. plt.subplot(1, 3, 2)
  295. plt.plot(history.history['auc'], label='Training AUC')
  296. plt.plot(history.history['val_auc'], label='Validation AUC')
  297. # Mark best epoch if provided
  298. if best_epoch is not None:
  299. plt.axvline(x=best_epoch, color='r', linestyle='--', alpha=0.7)
  300. plt.plot(best_epoch, history.history['val_auc'][best_epoch],
  301. 'ro', markersize=8)
  302. plt.title('Model AUC')
  303. plt.xlabel('Epochs')
  304. plt.ylabel('AUC')
  305. plt.legend()
  306. plt.grid(True, alpha=0.3)
  307. plt.ylim(0, 1) # Set y-axis from 0 to 1
  308. # Subplot 3: Accuracy (y-axis starts from 0)
  309. plt.subplot(1, 3, 3)
  310. plt.plot(history.history['accuracy'], label='Training Accuracy')
  311. plt.plot(history.history['val_accuracy'], label='Validation Accuracy')
  312. # Mark best epoch if provided
  313. if best_epoch is not None:
  314. plt.axvline(x=best_epoch, color='r', linestyle='--', alpha=0.7)
  315. plt.plot(best_epoch, history.history['val_accuracy'][best_epoch],
  316. 'ro', markersize=8, label=f'Best: {best_val_acc:.4f}')
  317. plt.title('Model Accuracy')
  318. plt.xlabel('Epochs')
  319. plt.ylabel('Accuracy')
  320. plt.legend()
  321. plt.grid(True, alpha=0.3)
  322. plt.ylim(0, 1) # Set y-axis from 0 to 1
  323. title_suffix = ""
  324. if best_epoch is not None:
  325. title_suffix = f" | Best Epoch: {best_epoch + 1}"
  326. plt.suptitle(f"Training History - Network Architecture: {params['layer_sizes']}{title_suffix}",
  327. fontsize=14)
  328. plt.tight_layout()
  329. plt.savefig('training_history_val_acc.png', dpi=Config.FIGURE_DPI, bbox_inches='tight')
  330. plt.show()
  331. # ==================== Model Evaluator ====================
  332. class ModelEvaluator:
  333. """Model evaluator"""
  334. @staticmethod
  335. def evaluate_model(model, X_train, y_train, X_val, y_val, model_name="Final Model", val_df=None):
  336. """Evaluate model performance and output statistical results"""
  337. print("\n" + "=" * 60)
  338. print(f"Model Performance Evaluation - {model_name}")
  339. print("=" * 60)
  340. # Training set predictions
  341. y_train_pred = model.predict(X_train, verbose=0)
  342. y_train_pred_class = (y_train_pred > 0.5).astype(int)
  343. y_train_pred = y_train_pred.flatten()
  344. # Validation set predictions
  345. y_val_pred = model.predict(X_val, verbose=0)
  346. y_val_pred_class = (y_val_pred > 0.5).astype(int)
  347. y_val_pred = y_val_pred.flatten()
  348. # Calculate metrics
  349. train_metrics = ModelEvaluator._calculate_metrics(y_train, y_train_pred, y_train_pred_class)
  350. val_metrics = ModelEvaluator._calculate_metrics(y_val, y_val_pred, y_val_pred_class)
  351. # Output performance metrics
  352. metrics_df = ModelEvaluator._create_metrics_table(train_metrics, val_metrics)
  353. # Output classification report
  354. ModelEvaluator._print_classification_report(y_val, y_val_pred_class)
  355. # Save validation set predictions if val_df is provided
  356. if val_df is not None:
  357. ModelEvaluator._save_validation_predictions(val_df, y_val, y_val_pred, y_val_pred_class, model_name)
  358. return metrics_df, y_train_pred, y_val_pred, y_val_pred_class
  359. @staticmethod
  360. def _calculate_metrics(y_true, y_pred_prob, y_pred_class):
  361. """Calculate all evaluation metrics"""
  362. return {
  363. 'Accuracy': accuracy_score(y_true, y_pred_class),
  364. 'Precision': precision_score(y_true, y_pred_class, zero_division=0),
  365. 'Recall': recall_score(y_true, y_pred_class, zero_division=0),
  366. 'F1-Score': f1_score(y_true, y_pred_class, zero_division=0),
  367. 'AUC': roc_auc_score(y_true, y_pred_prob),
  368. 'Specificity': ModelEvaluator._calculate_specificity(y_true, y_pred_class)
  369. }
  370. @staticmethod
  371. def _calculate_specificity(y_true, y_pred):
  372. """Calculate specificity"""
  373. tn, fp, fn, tp = confusion_matrix(y_true, y_pred).ravel()
  374. return tn / (tn + fp) if (tn + fp) > 0 else 0
  375. @staticmethod
  376. def _create_metrics_table(train_metrics, val_metrics):
  377. """Create performance metrics table"""
  378. metrics_df = pd.DataFrame({
  379. 'Training Set': train_metrics,
  380. 'Validation Set': val_metrics
  381. })
  382. print("\nPerformance Metrics Comparison:")
  383. print(metrics_df.round(4))
  384. # Calculate differences
  385. metrics_df['Difference (Validation-Training)'] = metrics_df['Validation Set'] - metrics_df['Training Set']
  386. print("\nTraining vs Validation Set Differences:")
  387. print(metrics_df['Difference (Validation-Training)'].round(4))
  388. return metrics_df
  389. @staticmethod
  390. def _print_classification_report(y_true, y_pred):
  391. """Print classification report"""
  392. print("\nValidation Set Classification Report:")
  393. print(classification_report(y_true, y_pred,
  394. target_names=['Non-toxic (Class 0)', 'Toxic (Class 1)']))
  395. @staticmethod
  396. def _save_validation_predictions(val_df, y_true, y_pred_prob, y_pred_class, model_name):
  397. """Save validation set predictions to CSV file"""
  398. try:
  399. # Create a DataFrame with predictions
  400. predictions_df = pd.DataFrame({
  401. 'smiles': val_df['smiles'].values,
  402. 'true_label': y_true,
  403. 'predicted_probability': y_pred_prob,
  404. 'predicted_label': y_pred_class,
  405. 'correct_prediction': (y_true == y_pred_class).astype(int)
  406. })
  407. # Add prediction confidence categories
  408. predictions_df['prediction_confidence'] = pd.cut(
  409. predictions_df['predicted_probability'],
  410. bins=[0, 0.3, 0.7, 1.0],
  411. labels=['Low', 'Medium', 'High'],
  412. include_lowest=True
  413. )
  414. # Sort by prediction probability (descending)
  415. predictions_df = predictions_df.sort_values('predicted_probability', ascending=False)
  416. # Create model-specific filename
  417. model_name_clean = model_name.lower().replace(' ', '_').replace('-', '_')
  418. filename = f'validation_predictions_{model_name_clean}.csv'
  419. # Save to CSV
  420. predictions_df.to_csv(filename, index=False)
  421. print(f"\n✓ Validation set predictions saved to: {filename}")
  422. print(f" - Total molecules: {len(predictions_df)}")
  423. print(f" - Correct predictions: {predictions_df['correct_prediction'].sum()} ({predictions_df['correct_prediction'].mean():.2%})")
  424. # Print summary statistics
  425. print(f"\n Prediction Statistics:")
  426. print(f" - Mean predicted probability: {predictions_df['predicted_probability'].mean():.4f}")
  427. print(f" - Std predicted probability: {predictions_df['predicted_probability'].std():.4f}")
  428. print(f" - Min predicted probability: {predictions_df['predicted_probability'].min():.4f}")
  429. print(f" - Max predicted probability: {predictions_df['predicted_probability'].max():.4f}")
  430. # Count by confidence level
  431. confidence_counts = predictions_df['prediction_confidence'].value_counts().sort_index()
  432. print(f"\n Confidence Level Distribution:")
  433. for conf_level, count in confidence_counts.items():
  434. percentage = count / len(predictions_df) * 100
  435. print(f" - {conf_level}: {count} molecules ({percentage:.1f}%)")
  436. except Exception as e:
  437. print(f"⚠ Warning: Failed to save validation predictions: {e}")
  438. @staticmethod
  439. def plot_comparison_visualizations(y_train, y_train_pred_prob_best, y_train_pred_prob_final,
  440. y_val, y_val_pred_prob_best, y_val_pred_prob_final,
  441. y_val_true, y_val_pred_class_best,
  442. train_metrics_best, val_metrics_best,
  443. train_metrics_final=None, val_metrics_final=None):
  444. """Plot comparison visualizations between best and final models"""
  445. fig = plt.figure(figsize=(16, 10))
  446. # Subplot 1: ROC Curves Comparison
  447. ax1 = plt.subplot(2, 3, 1)
  448. # Calculate ROC curves for best model
  449. fpr_train_best, tpr_train_best, _ = roc_curve(y_train, y_train_pred_prob_best)
  450. fpr_val_best, tpr_val_best, _ = roc_curve(y_val, y_val_pred_prob_best)
  451. ax1.plot(fpr_train_best, tpr_train_best, 'b-',
  452. label=f'Training (Best) AUC = {train_metrics_best["AUC"]:.3f}')
  453. ax1.plot(fpr_val_best, tpr_val_best, 'r-',
  454. label=f'Validation (Best) AUC = {val_metrics_best["AUC"]:.3f}')
  455. # Calculate ROC curves for final model if provided
  456. if y_train_pred_prob_final is not None and y_val_pred_prob_final is not None:
  457. fpr_train_final, tpr_train_final, _ = roc_curve(y_train, y_train_pred_prob_final)
  458. fpr_val_final, tpr_val_final, _ = roc_curve(y_val, y_val_pred_prob_final)
  459. ax1.plot(fpr_train_final, tpr_train_final, 'b--', alpha=0.7,
  460. label=f'Training (Final) AUC = {train_metrics_final["AUC"]:.3f}')
  461. ax1.plot(fpr_val_final, tpr_val_final, 'r--', alpha=0.7,
  462. label=f'Validation (Final) AUC = {val_metrics_final["AUC"]:.3f}')
  463. ax1.plot([0, 1], [0, 1], 'k--', label='Random Classifier')
  464. ax1.set_xlim([0.0, 1.0])
  465. ax1.set_ylim([0.0, 1.05])
  466. ax1.set_xlabel('False Positive Rate (FPR)')
  467. ax1.set_ylabel('True Positive Rate (TPR)')
  468. ax1.set_title('ROC Curve Comparison')
  469. ax1.legend(loc="lower right")
  470. ax1.grid(True, alpha=0.3)
  471. # Subplot 2: Confusion Matrix for Best Model
  472. ax2 = plt.subplot(2, 3, 2)
  473. cm_best = confusion_matrix(y_val_true, y_val_pred_class_best)
  474. sns.heatmap(cm_best, annot=True, fmt='d', cmap='Blues', ax=ax2,
  475. xticklabels=['Predicted Non-toxic', 'Predicted Toxic'],
  476. yticklabels=['Actual Non-toxic', 'Actual Toxic'])
  477. ax2.set_title('Confusion Matrix (Best Model)')
  478. ax2.set_ylabel('True Label')
  479. ax2.set_xlabel('Predicted Label')
  480. # Subplot 3: Metrics Comparison Bar Chart
  481. ax3 = plt.subplot(2, 3, 3)
  482. metrics_to_plot = ['Accuracy', 'Precision', 'Recall', 'F1-Score', 'AUC']
  483. x = np.arange(len(metrics_to_plot))
  484. width = 0.35
  485. train_values_best = [train_metrics_best[m] for m in metrics_to_plot]
  486. val_values_best = [val_metrics_best[m] for m in metrics_to_plot]
  487. bars1 = ax3.bar(x - width/2, train_values_best, width, label='Training (Best)', alpha=0.8, color='blue')
  488. bars2 = ax3.bar(x + width/2, val_values_best, width, label='Validation (Best)', alpha=0.8, color='red')
  489. # Add final model metrics if provided
  490. if train_metrics_final is not None and val_metrics_final is not None:
  491. train_values_final = [train_metrics_final[m] for m in metrics_to_plot]
  492. val_values_final = [val_metrics_final[m] for m in metrics_to_plot]
  493. bars3 = ax3.bar(x - width/2, train_values_final, width, label='Training (Final)',
  494. alpha=0.5, color='lightblue', hatch='//')
  495. bars4 = ax3.bar(x + width/2, val_values_final, width, label='Validation (Final)',
  496. alpha=0.5, color='lightcoral', hatch='\\\\')
  497. ax3.set_xlabel('Evaluation Metrics')
  498. ax3.set_ylabel('Score')
  499. ax3.set_title('Performance Metrics Comparison')
  500. ax3.set_xticks(x)
  501. ax3.set_xticklabels(metrics_to_plot, rotation=45)
  502. ax3.legend()
  503. ax3.grid(True, alpha=0.3, axis='y')
  504. ax3.set_ylim(0, 1) # Set y-axis from 0 to 1
  505. # Subplot 4: Validation Accuracy Comparison
  506. ax4 = plt.subplot(2, 3, 4)
  507. comparison_data = {
  508. 'Best Model': val_metrics_best['Accuracy'],
  509. }
  510. if val_metrics_final is not None:
  511. comparison_data['Final Epoch Model'] = val_metrics_final['Accuracy']
  512. models = list(comparison_data.keys())
  513. acc_values = list(comparison_data.values())
  514. bars = ax4.bar(models, acc_values, color=['green', 'orange'][:len(models)])
  515. ax4.set_ylabel('Validation Accuracy')
  516. ax4.set_title('Validation Accuracy Comparison')
  517. ax4.set_ylim([0, 1.0])
  518. # Add value labels on bars
  519. for i, (bar, val) in enumerate(zip(bars, acc_values)):
  520. height = bar.get_height()
  521. ax4.text(bar.get_x() + bar.get_width()/2., height + 0.01,
  522. f'{val:.4f}', ha='center', va='bottom')
  523. ax4.grid(True, alpha=0.3, axis='y')
  524. # Subplot 5: AUC Comparison
  525. ax5 = plt.subplot(2, 3, 5)
  526. auc_comparison_data = {
  527. 'Best Model': val_metrics_best['AUC'],
  528. }
  529. if val_metrics_final is not None:
  530. auc_comparison_data['Final Epoch Model'] = val_metrics_final['AUC']
  531. models_auc = list(auc_comparison_data.keys())
  532. auc_values = list(auc_comparison_data.values())
  533. bars_auc = ax5.bar(models_auc, auc_values, color=['green', 'orange'][:len(models_auc)])
  534. ax5.set_ylabel('AUC Score')
  535. ax5.set_title('AUC Score Comparison')
  536. ax5.set_ylim([0, 1.0])
  537. # Add value labels on bars
  538. for i, (bar, val) in enumerate(zip(bars_auc, auc_values)):
  539. height = bar.get_height()
  540. ax5.text(bar.get_x() + bar.get_width()/2., height + 0.01,
  541. f'{val:.4f}', ha='center', va='bottom')
  542. ax5.grid(True, alpha=0.3, axis='y')
  543. # Subplot 6: F1-Score Comparison
  544. ax6 = plt.subplot(2, 3, 6)
  545. f1_comparison_data = {
  546. 'Best Model': val_metrics_best['F1-Score'],
  547. }
  548. if val_metrics_final is not None:
  549. f1_comparison_data['Final Epoch Model'] = val_metrics_final['F1-Score']
  550. models_f1 = list(f1_comparison_data.keys())
  551. f1_values = list(f1_comparison_data.values())
  552. bars_f1 = ax6.bar(models_f1, f1_values, color=['green', 'orange'][:len(models_f1)])
  553. ax6.set_ylabel('F1-Score')
  554. ax6.set_title('F1-Score Comparison')
  555. ax6.set_ylim([0, 1.0])
  556. # Add value labels on bars
  557. for i, (bar, val) in enumerate(zip(bars_f1, f1_values)):
  558. height = bar.get_height()
  559. ax6.text(bar.get_x() + bar.get_width()/2., height + 0.01,
  560. f'{val:.4f}', ha='center', va='bottom')
  561. ax6.grid(True, alpha=0.3, axis='y')
  562. plt.suptitle('Model Performance Comparison: Best vs Final Epoch Model', fontsize=16)
  563. plt.tight_layout()
  564. plt.savefig('model_comparison.png', dpi=Config.FIGURE_DPI, bbox_inches='tight')
  565. plt.show()
  566. # ==================== Results Saver ====================
  567. class ResultSaver:
  568. """Results saver"""
  569. @staticmethod
  570. def save_results(best_model, scaler, best_metrics_df, history, best_epoch, best_val_acc,
  571. final_model=None, final_metrics_df=None, val_df=None,
  572. y_val=None, y_val_pred_best=None, y_val_pred_final=None):
  573. """Save all results"""
  574. print("\n" + "=" * 60)
  575. print("Saving Model and Results")
  576. print("=" * 60)
  577. # Create results directory
  578. os.makedirs('results_val_acc', exist_ok=True)
  579. # Get optimal parameters
  580. optimal_params = Config.OPTIMAL_PARAMS
  581. # 1. Save best model
  582. best_model.save('results_val_acc/best_model_val_acc.h5')
  583. print("✓ Best model saved: results_val_acc/best_model_val_acc.h5")
  584. # Save final model if provided
  585. if final_model is not None:
  586. final_model.save('results_val_acc/final_epoch_model.h5')
  587. print("✓ Final epoch model saved: results_val_acc/final_epoch_model.h5")
  588. # 2. Save scaler
  589. import joblib
  590. joblib.dump(scaler, 'results_val_acc/scaler.pkl')
  591. print("✓ Scaler saved: results_val_acc/scaler.pkl")
  592. # 3. Save validation set predictions for best model
  593. if val_df is not None and y_val is not None and y_val_pred_best is not None:
  594. ResultSaver._save_detailed_predictions(val_df, y_val, y_val_pred_best,
  595. 'best_model', optimal_params)
  596. # 4. Save validation set predictions for final model (optional)
  597. if val_df is not None and y_val is not None and y_val_pred_final is not None:
  598. ResultSaver._save_detailed_predictions(val_df, y_val, y_val_pred_final,
  599. 'final_epoch_model', optimal_params)
  600. # 5. Save optimal parameters and training info
  601. with open('results_val_acc/training_info.txt', 'w') as f:
  602. f.write("Optimal Model Parameters (based on cross-validation results):\n")
  603. f.write("=" * 50 + "\n")
  604. for param, value in optimal_params.items():
  605. f.write(f"{param}: {value}\n")
  606. f.write("\nModel Selection Information:\n")
  607. f.write("=" * 50 + "\n")
  608. f.write(f"Model Selection Metric: {Config.MODEL_SELECTION_METRIC}\n")
  609. f.write(f"Best Epoch: {best_epoch + 1}\n")
  610. f.write(f"Best Validation Accuracy: {best_val_acc:.4f}\n")
  611. f.write(f"Total Epochs Trained: {len(history.history['loss'])}\n")
  612. print("✓ Training information saved: results_val_acc/training_info.txt")
  613. # 6. Save performance metrics
  614. best_metrics_df.to_csv('results_val_acc/best_model_performance.csv')
  615. print("✓ Best model performance metrics saved: results_val_acc/best_model_performance.csv")
  616. if final_metrics_df is not None:
  617. final_metrics_df.to_csv('results_val_acc/final_epoch_performance.csv')
  618. print("✓ Final epoch performance metrics saved: results_val_acc/final_epoch_performance.csv")
  619. # 7. Save training history
  620. history_df = pd.DataFrame(history.history)
  621. history_df.to_csv('results_val_acc/training_history.csv')
  622. print("✓ Training history saved: results_val_acc/training_history.csv")
  623. # 8. Save configuration summary
  624. with open('results_val_acc/config_summary.txt', 'w') as f:
  625. f.write("Model Configuration Summary\n")
  626. f.write("=" * 60 + "\n\n")
  627. f.write("Overfitting Prevention Measures:\n")
  628. f.write(f" - L2 Regularization: {optimal_params['l2_reg']}\n")
  629. f.write(f" - Dropout Rate: {optimal_params['dropout_rate']}\n")
  630. f.write(f" - Learning Rate Reduction Factor: {Config.REDUCE_LR_FACTOR}\n")
  631. f.write(f" - Learning Rate Reduction Patience: {Config.REDUCE_LR_PATIENCE}\n")
  632. f.write(f" - Early Stopping Patience: {Config.EARLY_STOPPING_PATIENCE}\n\n")
  633. f.write("Network Architecture:\n")
  634. f.write(f" - Layer Structure: {optimal_params['layer_sizes']}\n")
  635. f.write(f" - Total Layers: {len(optimal_params['layer_sizes'])} (hidden) + 1 (output)\n\n")
  636. f.write("Training Parameters:\n")
  637. f.write(f" - Learning Rate: {optimal_params['learning_rate']}\n")
  638. f.write(f" - Batch Size: {optimal_params['batch_size']}\n")
  639. f.write(f" - Maximum Epochs: {optimal_params['epochs']}\n")
  640. f.write(f" - Actual Epochs Trained: {len(history.history['loss'])}\n")
  641. f.write(f" - Model Selection: Based on {Config.MODEL_SELECTION_METRIC}\n")
  642. f.write(f" - Best Model Epoch: {best_epoch + 1}\n\n")
  643. f.write("Dataset Information:\n")
  644. f.write(f" - Training Set Size: 2179\n")
  645. f.write(f" - Validation Set Size: 545\n")
  646. f.write(f" - Feature Dimension: 167 (MACCS fingerprints)\n")
  647. print("✓ Configuration summary saved: results_val_acc/config_summary.txt")
  648. # 9. Generate final report
  649. ResultSaver._generate_final_report(best_metrics_df, optimal_params, history,
  650. best_epoch, best_val_acc, final_metrics_df)
  651. @staticmethod
  652. def _save_detailed_predictions(val_df, y_true, y_pred_prob, model_type, optimal_params):
  653. """Save detailed validation set predictions for a specific model"""
  654. try:
  655. # Calculate predicted class
  656. y_pred_class = (y_pred_prob > 0.5).astype(int)
  657. # Create detailed predictions DataFrame
  658. predictions_df = pd.DataFrame({
  659. 'smiles': val_df['smiles'].values,
  660. 'true_label': y_true,
  661. 'predicted_probability': y_pred_prob,
  662. 'predicted_label': y_pred_class,
  663. 'prediction_correct': (y_true == y_pred_class).astype(int),
  664. 'prediction_error': np.abs(y_true - y_pred_prob)
  665. })
  666. # Add additional information if available in original dataframe
  667. for col in val_df.columns:
  668. if col not in ['smiles', 'label'] and col not in predictions_df.columns:
  669. predictions_df[col] = val_df[col].values
  670. # Add prediction confidence level
  671. predictions_df['confidence_level'] = pd.cut(
  672. predictions_df['predicted_probability'],
  673. bins=[0, 0.3, 0.7, 1.0],
  674. labels=['Low (0-0.3)', 'Medium (0.3-0.7)', 'High (0.7-1.0)'],
  675. include_lowest=True
  676. )
  677. # Calculate prediction performance metrics for each confidence level
  678. confidence_stats = []
  679. for level in ['Low (0-0.3)', 'Medium (0.3-0.7)', 'High (0.7-1.0)']:
  680. subset = predictions_df[predictions_df['confidence_level'] == level]
  681. if len(subset) > 0:
  682. accuracy = subset['prediction_correct'].mean()
  683. confidence_stats.append({
  684. 'confidence_level': level,
  685. 'n_molecules': len(subset),
  686. 'accuracy': accuracy,
  687. 'avg_probability': subset['predicted_probability'].mean()
  688. })
  689. # Sort by prediction probability (descending)
  690. predictions_df = predictions_df.sort_values('predicted_probability', ascending=False)
  691. # Save predictions to CSV
  692. filename = f'results_val_acc/{model_type}_validation_predictions.csv'
  693. predictions_df.to_csv(filename, index=False)
  694. # Save confidence level statistics
  695. if confidence_stats:
  696. confidence_df = pd.DataFrame(confidence_stats)
  697. confidence_filename = f'results_val_acc/{model_type}_confidence_statistics.csv'
  698. confidence_df.to_csv(confidence_filename, index=False)
  699. print(f"✓ {model_type} confidence statistics saved: {confidence_filename}")
  700. # Generate summary file
  701. summary_filename = f'results_val_acc/{model_type}_predictions_summary.txt'
  702. with open(summary_filename, 'w') as f:
  703. f.write(f"{'=' * 60}\n")
  704. f.write(f"{model_type.replace('_', ' ').title()} - Validation Set Predictions Summary\n")
  705. f.write(f"{'=' * 60}\n\n")
  706. f.write(f"Model Parameters:\n")
  707. f.write(f" - Network Architecture: {optimal_params['layer_sizes']}\n")
  708. f.write(f" - Dropout Rate: {optimal_params['dropout_rate']}\n")
  709. f.write(f" - Learning Rate: {optimal_params['learning_rate']}\n")
  710. f.write(f" - L2 Regularization: {optimal_params['l2_reg']}\n\n")
  711. f.write(f"Overall Performance:\n")
  712. f.write(f" - Total Molecules: {len(predictions_df)}\n")
  713. f.write(f" - Correct Predictions: {predictions_df['prediction_correct'].sum()} ({predictions_df['prediction_correct'].mean():.2%})\n")
  714. f.write(f" - Mean Predicted Probability: {predictions_df['predicted_probability'].mean():.4f}\n")
  715. f.write(f" - Std Predicted Probability: {predictions_df['predicted_probability'].std():.4f}\n\n")
  716. f.write(f"Confidence Level Distribution:\n")
  717. for stats in confidence_stats:
  718. f.write(f" - {stats['confidence_level']}: {stats['n_molecules']} molecules, Accuracy: {stats['accuracy']:.2%}, Avg Probability: {stats['avg_probability']:.4f}\n")
  719. f.write(f"\nPrediction Probability Statistics:\n")
  720. f.write(f" - Min: {predictions_df['predicted_probability'].min():.4f}\n")
  721. f.write(f" - 25th Percentile: {predictions_df['predicted_probability'].quantile(0.25):.4f}\n")
  722. f.write(f" - Median: {predictions_df['predicted_probability'].median():.4f}\n")
  723. f.write(f" - 75th Percentile: {predictions_df['predicted_probability'].quantile(0.75):.4f}\n")
  724. f.write(f" - Max: {predictions_df['predicted_probability'].max():.4f}\n\n")
  725. f.write(f"Misclassified Molecules (Top 10 by Error):\n")
  726. misclassified = predictions_df[predictions_df['prediction_correct'] == 0]
  727. if len(misclassified) > 0:
  728. top_misclassified = misclassified.nlargest(10, 'prediction_error')
  729. for idx, row in top_misclassified.iterrows():
  730. f.write(f" - SMILES: {row['smiles']}, True: {row['true_label']}, Pred: {row['predicted_label']} (Prob: {row['predicted_probability']:.4f})\n")
  731. else:
  732. f.write(f" - No misclassified molecules!\n")
  733. print(f"✓ {model_type} validation predictions saved: {filename}")
  734. print(f"✓ {model_type} predictions summary saved: {summary_filename}")
  735. except Exception as e:
  736. print(f"⚠ Warning: Failed to save {model_type} predictions: {e}")
  737. @staticmethod
  738. def _generate_final_report(best_metrics_df, optimal_params, history,
  739. best_epoch, best_val_acc, final_metrics_df=None):
  740. """Generate final report"""
  741. actual_epochs = len(history.history['loss'])
  742. report = f"""
  743. {'=' * 60}
  744. Developmental Neurotoxicity Prediction DNN Model - Training Complete Report
  745. {'=' * 60}
  746. Model Configuration:
  747. ----------
  748. Network Architecture: {optimal_params['layer_sizes']}
  749. Dropout Rate: {optimal_params['dropout_rate']}
  750. Learning Rate: {optimal_params['learning_rate']}
  751. L2 Regularization: {optimal_params['l2_reg']}
  752. Batch Size: {optimal_params['batch_size']}
  753. Maximum Epochs: {optimal_params['epochs']}
  754. Actual Epochs Trained: {actual_epochs}
  755. Model Selection: Based on {Config.MODEL_SELECTION_METRIC}
  756. Best Model Epoch: {best_epoch + 1}
  757. Performance Summary (Best Model):
  758. ----------
  759. Training Set AUC: {best_metrics_df.loc['AUC', 'Training Set']:.4f}
  760. Validation Set AUC: {best_metrics_df.loc['AUC', 'Validation Set']:.4f}
  761. Training Set Accuracy: {best_metrics_df.loc['Accuracy', 'Training Set']:.4f}
  762. Validation Set Accuracy: {best_metrics_df.loc['Accuracy', 'Validation Set']:.4f}
  763. (Best Validation Accuracy: {best_val_acc:.4f})
  764. Training Set F1-Score: {best_metrics_df.loc['F1-Score', 'Training Set']:.4f}
  765. Validation Set F1-Score: {best_metrics_df.loc['F1-Score', 'Validation Set']:.4f}
  766. Overfitting Assessment (Best Model):
  767. -----------------
  768. AUC Difference (Validation-Training): {best_metrics_df.loc['AUC', 'Difference (Validation-Training)']:.4f}
  769. F1-Score Difference (Validation-Training): {best_metrics_df.loc['F1-Score', 'Difference (Validation-Training)']:.4f}
  770. """
  771. if final_metrics_df is not None:
  772. report += f"""
  773. Performance Comparison (Final Epoch vs Best Model):
  774. -----------------
  775. Validation Accuracy:
  776. - Final Epoch: {final_metrics_df.loc['Accuracy', 'Validation Set']:.4f}
  777. - Best Model: {best_metrics_df.loc['Accuracy', 'Validation Set']:.4f}
  778. - Improvement: {best_metrics_df.loc['Accuracy', 'Validation Set'] - final_metrics_df.loc['Accuracy', 'Validation Set']:.4f}
  779. Validation AUC:
  780. - Final Epoch: {final_metrics_df.loc['AUC', 'Validation Set']:.4f}
  781. - Best Model: {best_metrics_df.loc['AUC', 'Validation Set']:.4f}
  782. - Improvement: {best_metrics_df.loc['AUC', 'Validation Set'] - final_metrics_df.loc['AUC', 'Validation Set']:.4f}
  783. """
  784. report += f"""
  785. Evaluation Results:
  786. ---------
  787. {'Excellent Validation Performance (AUC > 0.85)' if best_metrics_df.loc['AUC', 'Validation Set'] > 0.85 else 'Moderate Validation Performance'}
  788. {'No Significant Overfitting (|AUC Difference| < 0.1)' if abs(best_metrics_df.loc['AUC', 'Difference (Validation-Training)']) < 0.1 else 'Potential Overfitting'}
  789. Prediction Files Generated:
  790. -----------------
  791. 1. best_model_validation_predictions.csv - Detailed predictions for each validation molecule
  792. 2. best_model_confidence_statistics.csv - Performance by confidence level
  793. 3. best_model_predictions_summary.txt - Summary statistics
  794. Note: The model was selected based on validation accuracy. The best model (epoch {best_epoch + 1})
  795. has been saved and will be used for all subsequent analysis.
  796. All results have been saved to the 'results_val_acc' directory.
  797. {'=' * 60}
  798. """
  799. with open('results_val_acc/final_report.txt', 'w') as f:
  800. f.write(report)
  801. print("\n" + report)
  802. # ==================== Main Program ====================
  803. def check_dependencies():
  804. """Check if all required dependencies are installed"""
  805. required_packages = {
  806. 'numpy': 'np',
  807. 'pandas': 'pd',
  808. 'matplotlib': 'plt',
  809. 'seaborn': 'sns',
  810. 'sklearn': 'sklearn',
  811. 'rdkit': 'Chem',
  812. 'tensorflow': 'tf'
  813. }
  814. missing_packages = []
  815. for package, import_name in required_packages.items():
  816. try:
  817. if package == 'rdkit':
  818. __import__('rdkit.Chem')
  819. else:
  820. __import__(package)
  821. except ImportError:
  822. missing_packages.append(package)
  823. if missing_packages:
  824. print("Missing required dependencies:")
  825. for package in missing_packages:
  826. print(f" - {package}")
  827. print("\nPlease install using the following commands:")
  828. print("pip install numpy pandas matplotlib seaborn scikit-learn tensorflow")
  829. print("conda install -c conda-forge rdkit # or use conda to install rdkit")
  830. return False
  831. return True
  832. def setup_matplotlib():
  833. """Setup matplotlib style"""
  834. try:
  835. # Set seaborn style
  836. sns.set_style("whitegrid")
  837. sns.set_palette("husl")
  838. print("✓ Using seaborn style")
  839. except Exception as e:
  840. print(f"⚠ Style setup failed: {e}, using default style")
  841. def main():
  842. """Main function: Execute complete model training pipeline"""
  843. print("=" * 60)
  844. print("Developmental Neurotoxicity Prediction - DNN QSAR Model")
  845. print("=" * 60)
  846. print("Model Selection Strategy: Using validation accuracy to select best model")
  847. print("Training with Early Stopping: Enabled")
  848. print("=" * 60)
  849. # Check dependencies
  850. if not check_dependencies():
  851. return
  852. # Setup matplotlib style
  853. setup_matplotlib()
  854. # Initialize components
  855. config = Config()
  856. data_processor = DataProcessor()
  857. model_trainer = ModelTrainer(config)
  858. model_evaluator = ModelEvaluator()
  859. result_saver = ResultSaver()
  860. try:
  861. # Step 1: Load and preprocess data
  862. print("\n[Step 1/6] Loading and preprocessing data...")
  863. train_df, val_df = data_processor.load_data(
  864. config.TRAIN_FILE, config.VAL_FILE
  865. )
  866. X_train, y_train, X_val, y_val, scaler = data_processor.prepare_features(
  867. train_df, val_df
  868. )
  869. # Step 2: Display optimal parameters information
  870. print("\n[Step 2/6] Configuring model with optimal parameters...")
  871. optimal_params = config.OPTIMAL_PARAMS
  872. print(f"\nOptimal Parameter Configuration:")
  873. print(f" Network Architecture: {optimal_params['layer_sizes']}")
  874. print(f" Dropout Rate: {optimal_params['dropout_rate']}")
  875. print(f" Learning Rate: {optimal_params['learning_rate']}")
  876. print(f" L2 Regularization: {optimal_params['l2_reg']}")
  877. print(f" Batch Size: {optimal_params['batch_size']}")
  878. print(f" Maximum Epochs: {optimal_params['epochs']}")
  879. print(f" Early Stopping Patience: {config.EARLY_STOPPING_PATIENCE}")
  880. print(f" Model Selection Metric: {config.MODEL_SELECTION_METRIC}")
  881. # Step 3: Train model with validation accuracy selection and early stopping
  882. print("\n[Step 3/6] Training model with validation accuracy selection and early stopping...")
  883. best_model, history, final_model = model_trainer.train_final_model(
  884. X_train, y_train, X_val, y_val
  885. )
  886. # Step 4: Plot training history with best epoch marked and y-axis starting from 0
  887. print("\n[Step 4/6] Plotting training history (y-axis starts from 0)...")
  888. model_trainer.plot_training_history(history, optimal_params,
  889. model_trainer.best_epoch, model_trainer.best_val_accuracy)
  890. # Step 5: Evaluate both best model and final epoch model
  891. print("\n[Step 5/6] Evaluating model performance...")
  892. # Evaluate best model
  893. print("\n" + "=" * 50)
  894. print("Evaluating Best Model (Selected by Validation Accuracy)")
  895. print("=" * 50)
  896. best_metrics_df, y_train_pred_best, y_val_pred_best, y_val_pred_class_best = model_evaluator.evaluate_model(
  897. best_model, X_train, y_train, X_val, y_val, "Best Model", val_df
  898. )
  899. # Evaluate final epoch model
  900. print("\n" + "=" * 50)
  901. print("Evaluating Final Epoch Model")
  902. print("=" * 50)
  903. final_metrics_df, y_train_pred_final, y_val_pred_final, y_val_pred_class_final = model_evaluator.evaluate_model(
  904. final_model, X_train, y_train, X_val, y_val, "Final Epoch Model", val_df
  905. )
  906. # Step 6: Plot comparison visualizations
  907. print("\n[Step 6/6] Plotting performance comparison...")
  908. model_evaluator.plot_comparison_visualizations(
  909. y_train, y_train_pred_best, y_train_pred_final,
  910. y_val, y_val_pred_best, y_val_pred_final,
  911. y_val, y_val_pred_class_best,
  912. best_metrics_df.loc[:, 'Training Set'].to_dict(),
  913. best_metrics_df.loc[:, 'Validation Set'].to_dict(),
  914. final_metrics_df.loc[:, 'Training Set'].to_dict(),
  915. final_metrics_df.loc[:, 'Validation Set'].to_dict()
  916. )
  917. # Save all results
  918. result_saver.save_results(
  919. best_model, scaler, best_metrics_df, history,
  920. model_trainer.best_epoch, model_trainer.best_val_accuracy,
  921. final_model, final_metrics_df, val_df, y_val, y_val_pred_best, y_val_pred_final
  922. )
  923. print("\n" + "=" * 60)
  924. print("🎉 Model training and evaluation complete!")
  925. print("=" * 60)
  926. print(f"\nSummary:")
  927. print(f" - Best model saved from epoch: {model_trainer.best_epoch + 1}")
  928. print(f" - Best validation accuracy: {model_trainer.best_val_accuracy:.4f}")
  929. print(f" - Final epoch validation accuracy: {final_metrics_df.loc['Accuracy', 'Validation Set']:.4f}")
  930. print(f" - Improvement by model selection: {model_trainer.best_val_accuracy - final_metrics_df.loc['Accuracy', 'Validation Set']:.4f}")
  931. if model_trainer.early_stopped_epoch:
  932. print(f" - Early stopping triggered at epoch: {model_trainer.early_stopped_epoch}")
  933. print(f"\nPrediction Files Generated:")
  934. print(f" - results_val_acc/best_model_validation_predictions.csv")
  935. print(f" - results_val_acc/best_model_confidence_statistics.csv")
  936. print(f" - results_val_acc/best_model_predictions_summary.txt")
  937. print("=" * 60)
  938. except FileNotFoundError as e:
  939. print(f"\n❌ Error: File not found - {e}")
  940. print("Please ensure the following files exist in the current directory:")
  941. print(f" - {config.TRAIN_FILE}")
  942. print(f" - {config.VAL_FILE}")
  943. except Exception as e:
  944. print(f"\n❌ Error: {e}")
  945. import traceback
  946. traceback.print_exc()
  947. if __name__ == "__main__":
  948. # Check required files
  949. if not os.path.exists(Config.TRAIN_FILE):
  950. print(f"Error: Training file {Config.TRAIN_FILE} does not exist")
  951. print(f"Please ensure the following files exist in the current directory:")
  952. print(f"1. {Config.TRAIN_FILE}")
  953. print(f"2. {Config.VAL_FILE}")
  954. elif not os.path.exists(Config.VAL_FILE):
  955. print(f"Error: Validation file {Config.VAL_FILE} does not exist")
  956. else:
  957. # Run main program
  958. main()

DNN-test_selection.py at commit 514073b, no license · at the source

Overview

Authors: Hongting Ma1, Wenhui Zhang1, Fengxi Liu1, Rong Ni1, Xue Wang1, Yanying Sun1, Xuelin Sun2, Xiao Li1,3
ORCID iDs: Xiao Li
  1. Shandong Key Laboratory of Digital Diagnosis and Treatment of Thoracic Oncology, Shandong Engineering Research Center of Precision Diagnosis and Treatment Technology for Neuro-Oncology, Department of Clinical Pharmacy, The First Affiliated Hospital of Shandong First Medical University, Shandong Provincial Qianfoshan Hospital Jinan 250014 China
  2. Department of Pharmacy, Beijing Hospital, National Center of Gerontology, Institute of Geriatric Medicine, Chinese Academy of Medical Sciences Beijing 100730 China
  3. State Key Laboratory of Neurology and Oncology Drug Development, The First Affiliated Hospital of Shandong First Medical University, Shandong Provincial Qianfoshan Hospital Jinan 250014 China
Journal: RSC advances, volume 16, issue 36, pages 37450-37464
Dates: received 19 April 2026; accepted 24 June 2026; published online 6 July 2026
Type: Research article · Language: English
License: CC BY-NC
Identifiers: DOI 10.1039/d6ra03343a · PMID 42440932 · PMCID PMC13334446 · OpenAlex W7167493370
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: methods / tools (subfield)
Methods: Statistics, Smoothing, state filtering, decompositions, Machine learning
Journal subjects: Chemistry
Topic: Anesthesia and Neurotoxicity Research (Developmental Neuroscience, Neuroscience), according to OpenAlex
Citations: not cited yet (Europe PMC); 48 references in the paper

Abstract

Developmental neurotoxicity (DNT) represents a critical yet underevaluated toxicity endpoint within current chemical safety assessment frameworks, particularly in the context of drug development. In the present study, we constructed a standardized benchmark dataset comprising 2724 structurally curated compounds, including 1362 DNT-positive chemicals derived from the Developmental Neurotox Reference List (DNTREF) of the U.S. Environmental Protection Agency (EPA) and 1362 property-matched presumed negative compounds (non-neuroactive drug-like compounds) selected from the ChEMBL database. Using this benchmark dataset, we systematically evaluated seven representative quantitative structure activity relationship (QSAR) modeling paradigms, covering traditional machine learning, deep neural networks, molecular language models, and graph neural networks. Across all models, the area under the receiver operating characteristic curve (AUC) values ranged from 0.85 to 0.90 on an independent test set. Among these, a deep neural network based on MACCS fingerprints achieved the highest sensitivity in identifying DNT-positive compounds, supporting its utility in early-stage drug safety screening where minimizing false negatives is critical. Furthermore, SHAP-based interpretability analysis revealed key structural motifs associated with DNT, highlighting the roles of hydrophobic aromatic frameworks, electrophilic reactivity, and polar functional groups in modulating predicted neurodevelopmental risk. Collectively, this study provides a reproducible benchmark dataset and an interpretable QSAR framework that supports early DNT risk prioritization and offers actionable structural insights for medicinal chemistry optimization in drug discovery.

Reproduced under the paper's license (CC BY-NC), from the paper cited above.

Repository

Its files are read in the Code ↔ Paper reader above, with 13 matches between paragraphs and lines of code.

lixiao1688/DNT_Benchmark_Dataset_Model

License: none: the authors keep all their rights
State: the link answers, verified on 27 September 2026
Evidence: files inventoried
Commit: 514073b77308152898322ed1e5380f12c0f7b426, 27 January 2026
Languages: Python (8)
Size: 16 files, 8 scripts
Software Heritage: not archived
Found in: “Data availability”
Holds: README
Not found: license file, CITATION.cff, environment file, tests, continuous integration, documentation
Tools: pandas (8 files), RDKit (8 files), NumPy (7 files), Matplotlib (6 files), scikit-learn (5 files), seaborn (5 files), Keras (4 files), TensorFlow (4 files), UMAP (2 files), SciPy (1 file), SHAP (1 file)
Availability: 1 check, the latest on 27 September 2026: the link answers
  • 27 September 2026: the link answers
9 files

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 1 repository of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 8 scripts, each with its path and the digest of its content;
  • 13 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

No dataset and no data link were found in the paper.

Data and model availability

The developmental neurotoxicity benchmark dataset constructed in this study includes standardized SMILES representations of all compounds, corresponding binary toxicity labels, and the fixed training, validation, and independent test set partitions used for model development and evaluation. These data are deposited in a public repository to ensure full reproducibility of the reported results. The dataset is expected to serve as a valuable resource for the computational toxicology and drug discovery communities, enabling further methodological development and benchmarking.

In addition, the scripts for MACCS fingerprint generation, the trained weights of the optimal DNN_MACCS model, and all scripts used for model training, inference, and SHAP-based interpretability analysis will be made publicly available.

All scripts for AD analysis and chemical space visualization using UMAP and t-SNE will also be released to support transparent assessment of model scope and result interpretation.

Reproduced under the paper's license (CC BY-NC), from the paper cited above.

Data availability

The data utilized in this study were derived from the DNTREF from U.S. EPA, and ChEMBL database. Complete data (including processed datasets, and all scripts necessary to reproduce the study) is available on GitHub at https://github.com/lixiao1688/DNT_Benchmark_Dataset_Model.git.

Supplementary information (SI): detailed data processing for benchmark dataset; detailed network architectures, hyperparameter settings, and training procedures; a summary of predictive performance for all models on validation set; unified chemical standardization process for data collection, structure preprocessing, and sample selection; training history and model comparison of the DNN_MACCS model. See DOI: https://doi.org/10.1039/d6ra03343a.

Reproduced under the paper's license (CC BY-NC), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 27 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 8 authors, 4 funders, 46 references.

Cite

This paper

Ma, H., Zhang, W., Liu, F., Ni, R., Wang, X., Sun, Y., Sun, X., & Li, X. (2026). A benchmark dataset and interpretable deep learning framework for drug-induced developmental neurotoxicity prediction. RSC advances, 16(36), 37450-37464. https://doi.org/10.1039/d6ra03343a

BibTeX

@article{ma2026benchmark,
author = {Ma, Hongting and Zhang, Wenhui and Liu, Fengxi and Ni, Rong and Wang, Xue and Sun, Yanying and Sun, Xuelin and Li, Xiao},
title = {{A benchmark dataset and interpretable deep learning framework for drug-induced developmental neurotoxicity prediction}},
journal = {RSC advances},
year = {2026},
month = jul,
volume = {16},
number = {36},
pages = {37450--37464},
publisher = {Royal Society of Chemistry},
issn = {2046-2069},
doi = {10.1039/d6ra03343a},
url = {https://doi.org/10.1039/d6ra03343a},
pmid = {42440932},
pmcid = {PMC13334446}
}

RIS

TY - JOUR
AU - Ma, Hongting
AU - Zhang, Wenhui
AU - Liu, Fengxi
AU - Ni, Rong
AU - Wang, Xue
AU - Sun, Yanying
AU - Sun, Xuelin
AU - Li, Xiao
TI - A benchmark dataset and interpretable deep learning framework for drug-induced developmental neurotoxicity prediction
T2 - RSC advances
J2 - RSC Adv
PY - 2026
DA - 2026/07/06
VL - 16
IS - 36
SP - 37450
EP - 37464
SN - 2046-2069
PB - Royal Society of Chemistry
DO - 10.1039/d6ra03343a
UR - https://doi.org/10.1039/d6ra03343a
LA - en
ER -

CSL-JSON

{
"id": "10.1039/d6ra03343a",
"type": "article-journal",
"title": "A benchmark dataset and interpretable deep learning framework for drug-induced developmental neurotoxicity prediction",
"container-title": "RSC advances",
"author": [
{
"family": "Ma",
"given": "Hongting"
},
{
"family": "Zhang",
"given": "Wenhui"
},
{
"family": "Liu",
"given": "Fengxi"
},
{
"family": "Ni",
"given": "Rong"
},
{
"family": "Wang",
"given": "Xue"
},
{
"family": "Sun",
"given": "Yanying"
},
{
"family": "Sun",
"given": "Xuelin"
},
{
"family": "Li",
"given": "Xiao"
}
],
"container-title-short": "RSC Adv",
"volume": "16",
"issue": "36",
"page": "37450-37464",
"DOI": "10.1039/d6ra03343a",
"PMID": "42440932",
"PMCID": "PMC13334446",
"ISSN": "2046-2069",
"publisher": "Royal Society of Chemistry",
"URL": "https://doi.org/10.1039/d6ra03343a",
"language": "en",
"issued": {
"date-parts": [
[
2026,
7,
6
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: RDKit, SHAP, Keras, 8 other tools
[2] doi:10.1093/nar/gkag706 [code]
scDifformer: diffusion-based post-training for virtual cell modeling across large-scale single-cell data.
Journal: Nucleic acids research
In common: RDKit, Keras, UMAP, 7 other tools, 1 reference
[3] doi:10.1038/s41598-026-48613-0 [code]
An snRNA-seq aging clock for the fruit fly head sheds light on sex-biased aging.
Journal: Scientific reports
In common: SHAP, Keras, TensorFlow, 6 other tools, 2 references
[4] doi:10.1126/sciadv.aed3650 [code]
Truthful visualizations for mass spectrometry imaging enable high-spatial-resolution interactive &lt;i&gt;m/z&lt;/i&gt; mapping and exploration.
Journal: Science advances
In common: SHAP, Keras, UMAP, 6 other tools, methods / tools
[5] doi:10.1177/13872877261453512 [code]
Fluorescence spectroscopy and machine learning methods for detection of Alzheimer's disease from circulating white blood cells.
Journal: Journal of Alzheimer's disease : JAD
In common: Keras, UMAP, TensorFlow, 6 other tools, methods / tools
[6] doi:10.1038/s41592-026-03057-2 [code]
CREsted: modeling genomic and synthetic cell-type-specific enhancers across tissues and species.
Journal: Nature methods
In common: Keras, UMAP, TensorFlow, 6 other tools, methods / tools
[7] doi:10.3390/s26175327 [code]
Subject Identity Confounds qEEG Emotion Recognition on DEAP and DREAMER.
Journal: Sensors (Basel, Switzerland)
In common: SHAP, Keras, TensorFlow, 6 other tools
[8] doi:10.1371/journal.pcbi.1014615 [code]
Toward reliable machine learning models for neural circuit inference: A diagnostic study of CNNs on spike trains.
Journal: PLoS computational biology
In common: SHAP, Keras, TensorFlow, 6 other tools
[9] doi:10.1038/s41467-026-75700-7 [code]
Gene regulatory innovations from transposable elements in primate cerebellum development.
Journal: Nature communications
In common: SHAP, Keras, TensorFlow, 6 other tools
[10] doi:10.1038/s41598-026-47047-y [code]
Imaging-genetics-based dementia risk prediction using deep survival neural networks in the Rotterdam Study.
Journal: Scientific reports
In common: SHAP, Keras, TensorFlow, 6 other tools

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.