OSCR

An snRNA-seq aging clock for the fruit fly head sheds light on sex-biased aging.

Code ↔ Paper

4 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 4 matches
  1. [1] § Methods › Model development and interpretation ↔ src/timeflies/models/model.py, lines 637–753 · score 0.83 · fully connected layers, convolutional neural network, convolution blocks, optimize, dense, flattening
  2. [2] § Methods › Model development and interpretation ↔ src/timeflies/models/model.py, lines 1230–1273 · score 0.77 · logistic regression, XGBoost, scikit learn, learning models, RandomForest
  3. [3] § Methods › Dataset ↔ src/timeflies/analysis/visuals.py, lines 187–251 · score 0.59 · sparse matrix, dense matrix, memory, zero
  4. [4] § Methods › Model development and interpretation ↔ src/timeflies/evaluation/interpreter.py, lines 191–263 · score 0.56 · GradientExplainer, neural networks, GPU, models

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,434 lines · 57 KB · CC-BY-NC-ND-4.0 · 2 matches

  1. import json
  2. import os
  3. import sys
  4. import dill as pickle
  5. import numpy as np
  6. import xgboost as xgb
  7. from sklearn.ensemble import RandomForestClassifier
  8. from sklearn.linear_model import LogisticRegression
  9. from sklearn.metrics import accuracy_score
  10. from sklearn.model_selection import train_test_split
  11. from ..utils.gpu_handler import suppress_stderr
  12. from ..utils.path_manager import PathManager
  13. # Import TensorFlow and related modules with suppressed stderr
  14. with suppress_stderr():
  15. import tensorflow as tf
  16. from tensorflow.keras.callbacks import EarlyStopping
  17. class CustomModelCheckpoint(tf.keras.callbacks.ModelCheckpoint):
  18. """
  19. CustomModelCheckpoint is a custom callback for saving model checkpoints.
  20. It inherits from the tf.keras.callbacks.ModelCheckpoint class and overrides some of its methods
  21. for saving model weights. In addition to the normal functionality, it also saves the best validation
  22. loss to a separate file and saves the model history to a separate file. This allows it to compare between an
  23. already saved model and a new model during training and save the new model only if it has a better validation loss.
  24. """
  25. def __init__(
  26. self,
  27. filepath,
  28. best_val_loss_path,
  29. label_path,
  30. label_encoder,
  31. reference_path,
  32. reference,
  33. scaler,
  34. scaler_path,
  35. is_scaler_fit,
  36. is_scaler_fit_path,
  37. highly_variable_genes,
  38. highly_variable_genes_path,
  39. mix_included,
  40. mix_included_path,
  41. num_features,
  42. num_features_path,
  43. metadata_path,
  44. path_manager=None,
  45. *args,
  46. **kwargs,
  47. ):
  48. """
  49. Initialize a CustomModelCheckpoint instance.
  50. Args:
  51. filepath (str): Path for saving the model weights.
  52. best_val_loss_path (str): Path for saving the best validation loss.
  53. label_path (str): Path for saving the label encoder.
  54. label_encoder (sklearn.preprocessing.LabelEncoder): The label encoder.
  55. reference_path (str): Path for saving the reference data.
  56. reference (numpy.ndarray): The reference data.
  57. scaler (sklearn.preprocessing.StandardScaler): The scaler.
  58. scaler_path (str): Path for saving the scaler.
  59. is_scaler_fit (bool): Whether the scaler has been fit or not.
  60. is_scaler_fit_path (str): Path for saving the is_scaler_fit variable.
  61. highly_variable_genes (list): List of highly variable genes.
  62. highly_variable_genes_path (str): Path for saving the highly variable genes.
  63. mix_included (bool): Whether mix is included.
  64. mix_included_path (str): Path for saving the mix_included variable.
  65. num_features (int): Number of features in the training data.
  66. num_features_path (str): Path to save the num_features variable.
  67. metadata_path (str): Path to the experiment metadata.json file.
  68. *args: Variable length argument list.
  69. **kwargs: Arbitrary keyword arguments.
  70. """
  71. # Call the parent class's constructor
  72. super().__init__(filepath, *args, **kwargs)
  73. # Initialize arguments
  74. self.best_val_loss = float("inf")
  75. self.best_val_loss_path = best_val_loss_path
  76. self.label_path = label_path
  77. self.label_encoder = label_encoder
  78. self.reference_path = reference_path
  79. self.reference = reference
  80. self.scaler = scaler
  81. self.scaler_path = scaler_path
  82. self.is_scaler_fit = is_scaler_fit
  83. self.is_scaler_fit_path = is_scaler_fit_path
  84. self.mix_included = mix_included
  85. self.mix_included_path = mix_included_path
  86. self.highly_variable_genes = highly_variable_genes
  87. self.highly_variable_genes_path = highly_variable_genes_path
  88. self.num_features = num_features
  89. self.num_features_path = num_features_path
  90. self.metadata_path = metadata_path
  91. # Track initial and current best validation losses
  92. self.initial_best_val_loss = float("inf") # The historical best before training
  93. self.model_improved = False # True only if final model beats historical best
  94. # Store path_manager for models/ folder saving
  95. self.path_manager = path_manager
  96. def set_best_val_loss(self, best_val_loss):
  97. """
  98. Set the best validation loss.
  99. Args:
  100. best_val_loss (float): The best validation loss.
  101. """
  102. # Store the historical best before training starts
  103. self.initial_best_val_loss = best_val_loss
  104. # Set the current best validation loss
  105. self.best_val_loss = best_val_loss
  106. self.best = best_val_loss # Update parent class's best variable
  107. def on_epoch_end(self, epoch, logs=None):
  108. """
  109. Method called at the end of an epoch during model's training. It checks if the current validation
  110. loss is better than the best validation loss seen so far and if so, saves the new best validation
  111. loss and calls the parent class's on_epoch_end method to save the model weights.
  112. Args:
  113. epoch (int): The number of the epoch that just finished.
  114. logs (dict, optional): Dictionary of logs, contains the metrics results for this training epoch.
  115. """
  116. # Get the current validation loss
  117. current_val_loss = logs.get("val_loss") if logs else None
  118. if current_val_loss is None:
  119. return
  120. # If the current validation loss is better than the best validation loss seen so far
  121. if float(current_val_loss) < self.best_val_loss:
  122. self.best_val_loss = float(current_val_loss)
  123. # Only set model_improved if this beats the historical best from before training
  124. if float(current_val_loss) < self.initial_best_val_loss:
  125. self.model_improved = True
  126. # Custom clean message instead of verbose Keras output
  127. print(
  128. f"\nEpoch {epoch + 1}: val_loss improved from {self.best:.5f} to {current_val_loss:.5f}"
  129. )
  130. # Save best validation loss to a file
  131. with open(self.best_val_loss_path, "w") as f:
  132. json.dump({"best_val_loss": self.best_val_loss}, f)
  133. # Save the label encoder (if exists - None for regression)
  134. with open(self.label_path, "wb") as label_file:
  135. pickle.dump(
  136. self.label_encoder, label_file
  137. ) # save the label_encoder when the model improves (None for regression)
  138. # Save reference data
  139. with open(self.reference_path, "wb") as reference_file:
  140. np.save(reference_file, self.reference)
  141. # Save scaler
  142. with open(self.scaler_path, "wb") as scaler_file:
  143. pickle.dump(self.scaler, scaler_file)
  144. # Save is_scaler_fit
  145. with open(self.is_scaler_fit_path, "wb") as is_scaler_fit_file:
  146. pickle.dump(self.is_scaler_fit, is_scaler_fit_file)
  147. # Save highly variable genes
  148. with open(
  149. self.highly_variable_genes_path, "wb"
  150. ) as highly_variable_genes_file:
  151. pickle.dump(self.highly_variable_genes, highly_variable_genes_file)
  152. with open(self.mix_included_path, "wb") as mix_included_file:
  153. pickle.dump(self.mix_included, mix_included_file)
  154. # Save num_features
  155. with open(self.num_features_path, "wb") as f:
  156. pickle.dump(self.num_features, f)
  157. # Call the parent class's on_epoch_end method to save the model weights
  158. # Temporarily suppress parent's verbose output
  159. original_verbose = self.verbose
  160. self.verbose = 0
  161. super().on_epoch_end(epoch, logs)
  162. self.verbose = original_verbose
  163. # Also save to models/ folder for reuse across evaluations
  164. if self.path_manager and self.model_improved:
  165. self._save_to_models_folder()
  166. def _save_to_models_folder(self):
  167. """Save model artifacts to models/ folder for reuse across evaluations."""
  168. import shutil
  169. from pathlib import Path
  170. # Get models folder path
  171. models_dir = Path(self.path_manager.get_models_folder_path())
  172. models_dir.mkdir(parents=True, exist_ok=True)
  173. # Copy all model artifacts to models/ folder
  174. artifacts = [
  175. (self.filepath, models_dir / Path(self.filepath).name),
  176. (self.best_val_loss_path, models_dir / "best_val_loss.json"),
  177. (self.label_path, models_dir / "label_encoder.pkl"),
  178. (self.reference_path, models_dir / "reference_data.npy"),
  179. (self.scaler_path, models_dir / "scaler.pkl"),
  180. (self.is_scaler_fit_path, models_dir / "is_scaler_fit.pkl"),
  181. (self.highly_variable_genes_path, models_dir / "highly_variable_genes.pkl"),
  182. (self.mix_included_path, models_dir / "mix_included.pkl"),
  183. (self.num_features_path, models_dir / "num_features.pkl"),
  184. ]
  185. for source, dest in artifacts:
  186. source_path = Path(source).resolve()
  187. dest_path = Path(dest).resolve()
  188. if source_path.exists() and source_path != dest_path:
  189. shutil.copy2(source_path, dest_path)
  190. class ModelLoader:
  191. """
  192. A class to manage the loading of machine learning or deep learning models.
  193. This class constructs the full path to the saved model based on the configuration settings
  194. and loads the model along with associated components, such as label encoders and other
  195. necessary preprocessing objects, to ensure the model is ready for inference.
  196. Attributes:
  197. - config (ConfigHandler): Holds configuration settings for loading the model and related components.
  198. - model_dir (str): Directory path where the model files are stored, constructed from configuration settings.
  199. - model_path (str): Full path to the specific model file to be loaded.
  200. - model_type (str): The type of model to be loaded (e.g., "cnn", "rnn"), specified in the configuration.
  201. """
  202. def __init__(
  203. self,
  204. config,
  205. pipeline_mode="training",
  206. use_models_folder=None,
  207. ):
  208. """
  209. Initializes the ModelLoader with configuration and directory structure to locate the model files.
  210. Parameters:
  211. - config (ConfigHandler): A ConfigHandler instance containing settings for model loading, paths,
  212. and preprocessing components.
  213. - pipeline_mode (str): "training" (train+eval) or "evaluation" (eval-only)
  214. - use_models_folder (bool): If True, load from models/ folder; if False, use best experiment;
  215. if None, auto-detect based on pipeline_mode
  216. Sets up:
  217. - `model_dir` by constructing the directory path from config details.
  218. - `model_path` as the specific file path for the saved model.
  219. - `model_type` to specify the type of model (e.g., CNN, RNN) as per config.
  220. """
  221. self.config = config
  222. self.pipeline_mode = pipeline_mode
  223. self.path_manager = PathManager(self.config)
  224. self.model_type = getattr(self.config.data, "model", "CNN").lower()
  225. # Determine whether to use models/ folder or best experiment
  226. if use_models_folder is None:
  227. # Auto-detect: use models/ for evaluation, best experiment for training
  228. use_models_folder = pipeline_mode == "evaluation"
  229. if use_models_folder:
  230. # Use models/ folder for trained model (evaluation mode)
  231. self.model_dir = self.path_manager.get_models_folder_path()
  232. else:
  233. # Use best experiment for current config (training mode)
  234. self.model_dir = self.path_manager.get_best_model_dir_for_config()
  235. self.model_path = self._get_model_path()
  236. def _get_model_path(self):
  237. """
  238. Determines the model file path based on the model type.
  239. Returns:
  240. - str: The path to the model file.
  241. """
  242. if self.model_type in ["cnn", "mlp"]:
  243. model_filename = "model.keras"
  244. else:
  245. model_filename = "model.pkl"
  246. model_path = os.path.join(self.model_dir, model_filename)
  247. return model_path
  248. def _verify_split_compatibility(self):
  249. """
  250. Verify that the current config's split settings are compatible with the saved model.
  251. Raises a warning or error if there's a mismatch that could cause issues.
  252. """
  253. import json
  254. from ..utils.split_naming import SplitNamingUtils
  255. # Look for metadata.json in the model directory
  256. metadata_path = os.path.join(self.model_dir, "metadata.json")
  257. if not os.path.exists(metadata_path):
  258. print(
  259. "WARNING: No metadata found for saved model, cannot verify split compatibility"
  260. )
  261. return
  262. try:
  263. with open(metadata_path) as f:
  264. saved_metadata = json.load(f)
  265. # Extract current configuration
  266. current_split = SplitNamingUtils.extract_split_details_for_metadata(
  267. self.config
  268. )
  269. # Check all key metadata fields for compatibility
  270. mismatches = []
  271. # Check model type
  272. current_model = getattr(self.config.model, "model_type", "CNN").upper()
  273. saved_model = saved_metadata.get("model_type", "")
  274. if current_model != saved_model:
  275. mismatches.append(f"Model type: {current_model} vs {saved_model}")
  276. # Check target variable
  277. current_target = getattr(self.config.data, "target_variable", "age")
  278. saved_target = saved_metadata.get("target", "")
  279. if current_target != saved_target:
  280. mismatches.append(f"Target: {current_target} vs {saved_target}")
  281. # Check tissue
  282. current_tissue = getattr(self.config.data, "tissue", "head")
  283. saved_tissue = saved_metadata.get("tissue", "")
  284. if current_tissue != saved_tissue:
  285. mismatches.append(f"Tissue: {current_tissue} vs {saved_tissue}")
  286. # Check batch correction
  287. current_batch = getattr(self.config.data.batch_correction, "enabled", False)
  288. saved_batch = saved_metadata.get("batch_correction", None)
  289. if saved_batch is not None and current_batch != saved_batch:
  290. mismatches.append(f"Batch correction: {current_batch} vs {saved_batch}")
  291. # Check split configuration
  292. saved_split = saved_metadata.get("split_config", {})
  293. if not saved_split:
  294. mismatches.append("Split configuration missing in saved model")
  295. else:
  296. method_match = current_split.get("method") == saved_split.get("method")
  297. split_name_match = current_split.get("split_name") == saved_split.get(
  298. "split_name"
  299. )
  300. if not method_match:
  301. mismatches.append(
  302. f"Split method: {current_split.get('method')} vs {saved_split.get('method')}"
  303. )
  304. if not split_name_match:
  305. mismatches.append(
  306. f"Split name: {current_split.get('split_name')} vs {saved_split.get('split_name')}"
  307. )
  308. # Report results
  309. if mismatches:
  310. print("WARNING: Configuration mismatches detected with saved model!")
  311. for mismatch in mismatches:
  312. print(f" - {mismatch}")
  313. print(
  314. " This may cause evaluation issues if configuration differs from model training"
  315. )
  316. elif self.pipeline_mode == "evaluation":
  317. # Only show success message during evaluation-only pipeline
  318. print("✓ Configuration matches saved model")
  319. except Exception as e:
  320. print(f"WARNING: Could not verify split compatibility: {e}")
  321. def load_model(self):
  322. """
  323. Loads the saved model and related components.
  324. This method constructs the file paths for the model and other related files
  325. (like label encoder, scaler, etc.) based on the configuration settings.
  326. It then loads these components and returns them for use.
  327. Returns:
  328. - tuple: A tuple containing the loaded model and related components like label encoder, reference data, scaler, test data, test labels, and training history.
  329. """
  330. # Verify split compatibility before loading
  331. self._verify_split_compatibility()
  332. # Load the model
  333. if os.path.exists(self.model_path):
  334. if self.model_type in ["cnn", "mlp"]:
  335. # Suppress the compile warning for loaded models
  336. import logging
  337. # Temporarily suppress absl warnings
  338. absl_logger = logging.getLogger("absl")
  339. old_level = absl_logger.level
  340. absl_logger.setLevel(logging.ERROR)
  341. model = tf.keras.models.load_model(self.model_path)
  342. # Restore logging level
  343. absl_logger.setLevel(old_level)
  344. else:
  345. model = self._load_pickle(self.model_path)
  346. else:
  347. print(f"ERROR: Model file not found: {self.model_path}")
  348. sys.exit(1)
  349. # Return all loaded components
  350. return model
  351. def load_model_components(self):
  352. """
  353. Loads the saved model's related components.
  354. This method constructs the file paths for the model and other related files
  355. (like label encoder, scaler, etc.) based on the configuration settings.
  356. It then loads these components and returns them for use.
  357. Returns:
  358. - tuple: A tuple containing the loaded model and related components like label encoder, reference data, scaler, test data, test labels, and training history.
  359. """
  360. # Load other related components from the model directory
  361. label_encoder = self._load_component_file("label_encoder.pkl")
  362. scaler = self._load_component_file("scaler.pkl")
  363. is_scaler_fit = self._load_component_file("is_scaler_fit.pkl")
  364. highly_variable_genes = self._load_component_file("highly_variable_genes.pkl")
  365. num_features = self._load_component_file("num_features.pkl")
  366. mix_included = self._load_component_file("mix_included.pkl")
  367. reference_data = self._load_component_file(
  368. "reference_data.npy", file_type="numpy"
  369. )
  370. # History is in training subdirectory
  371. history = self._load_training_file("history.pkl")
  372. # Return all loaded components
  373. return (
  374. label_encoder,
  375. scaler,
  376. is_scaler_fit,
  377. highly_variable_genes,
  378. num_features,
  379. history,
  380. mix_included,
  381. reference_data,
  382. )
  383. def _load_pickle(self, file_path):
  384. """
  385. Loads a pickle file from the given path.
  386. Parameters:
  387. - file_path (str): The path to the pickle file.
  388. Returns:
  389. - object: The object loaded from the pickle file.
  390. """
  391. with open(file_path, "rb") as file:
  392. return pickle.load(file)
  393. def _load_file(self, file_name, file_type="pickle"):
  394. """
  395. Loads a file from the model directory based on the file type.
  396. Parameters:
  397. - file_name (str): The name of the file to be loaded.
  398. - file_type (str): The type of the file to be loaded, either 'pickle' or 'numpy'.
  399. Returns:
  400. - object: The object loaded from the file, depending on its type (pickle or numpy).
  401. """
  402. file_path = os.path.join(self.model_dir, file_name)
  403. if os.path.exists(file_path):
  404. if file_type == "pickle":
  405. return self._load_pickle(file_path)
  406. elif file_type == "numpy":
  407. return np.load(file_path, allow_pickle=True)
  408. else:
  409. # Print error if the file does not exist and exit the program
  410. print(f"Error: {file_name} not found in {file_path}")
  411. sys.exit(1)
  412. def _load_component_file(self, file_name, file_type="pickle"):
  413. """
  414. Loads a component file, checking both new model_components/ and old root directory.
  415. Parameters:
  416. - file_name (str): The name of the file to be loaded.
  417. - file_type (str): The type of the file to be loaded, either 'pickle' or 'numpy'.
  418. Returns:
  419. - object: The object loaded from the file.
  420. """
  421. # Try new model_components directory first
  422. components_dir = os.path.join(self.model_dir, "model_components")
  423. new_path = os.path.join(components_dir, file_name)
  424. if os.path.exists(new_path):
  425. if file_type == "pickle":
  426. return self._load_pickle(new_path)
  427. elif file_type == "numpy":
  428. return np.load(new_path, allow_pickle=True)
  429. # Fallback to old location (root of model directory)
  430. old_path = os.path.join(self.model_dir, file_name)
  431. if os.path.exists(old_path):
  432. if file_type == "pickle":
  433. return self._load_pickle(old_path)
  434. elif file_type == "numpy":
  435. return np.load(old_path, allow_pickle=True)
  436. # File not found in either location
  437. print(f"Error: {file_name} not found in {new_path} or {old_path}")
  438. sys.exit(1)
  439. def _load_training_file(self, file_name, file_type="pickle"):
  440. """
  441. Loads a training file from models/ folder first, then fallback to training/ subdirectory or root.
  442. Parameters:
  443. - file_name (str): The name of the file to be loaded.
  444. - file_type (str): The type of the file to be loaded.
  445. Returns:
  446. - object: The object loaded from the file.
  447. """
  448. # For history.pkl, try models/ folder first (new location)
  449. if file_name == "history.pkl":
  450. models_path = os.path.join(
  451. self.path_manager.get_models_folder_path(), file_name
  452. )
  453. if os.path.exists(models_path):
  454. if file_type == "pickle":
  455. return self._load_pickle(models_path)
  456. elif file_type == "numpy":
  457. return np.load(models_path, allow_pickle=True)
  458. # Try training directory (old location for history)
  459. training_dir = os.path.join(self.model_dir, "training")
  460. new_path = os.path.join(training_dir, file_name)
  461. if os.path.exists(new_path):
  462. if file_type == "pickle":
  463. return self._load_pickle(new_path)
  464. elif file_type == "numpy":
  465. return np.load(new_path, allow_pickle=True)
  466. # Fallback to old location (root of model directory)
  467. old_path = os.path.join(self.model_dir, file_name)
  468. if os.path.exists(old_path):
  469. if file_type == "pickle":
  470. return self._load_pickle(old_path)
  471. elif file_type == "numpy":
  472. return np.load(old_path, allow_pickle=True)
  473. # File not found in any location
  474. if file_name == "history.pkl":
  475. print(
  476. f"Error: {file_name} not found in {self.path_manager.get_models_folder_path()}, {new_path} or {old_path}"
  477. )
  478. else:
  479. print(f"Error: {file_name} not found in {new_path} or {old_path}")
  480. sys.exit(1)
  481. class ModelBuilder:
  482. """
  483. A class to handle model building and training.
  484. This class constructs and trains a model based on the provided training data
  485. and configuration settings.
  486. """
  487. def __init__(
  488. self,
  489. config,
  490. train_data,
  491. train_labels,
  492. label_encoder,
  493. reference_data,
  494. scaler,
  495. is_scaler_fit,
  496. highly_variable_genes,
  497. mix_included,
  498. experiment_name=None,
  499. ):
  500. """
  501. Initializes the ModelBuilder with the given configuration and training data.
  502. Parameters:
  503. - config (ConfigHandler): A ConfigHandler object containing configuration settings.
  504. - train_data (numpy.ndarray): The training data.
  505. - train_labels (numpy.ndarray): The labels for the training data.
  506. - label_encoder (LabelEncoder): The label encoder for encoding labels.
  507. - reference_data (numpy.ndarray): Reference data used in model training.
  508. - scaler (object): Scaler object used for feature scaling.
  509. - is_scaler_fit (bool): Flag indicating if the scaler has been fitted.
  510. - highly_variable_genes (list): List of highly variable genes used in training.
  511. - mix_included (bool): Flag indicating if mix_included feature is used.
  512. """
  513. self.config = config
  514. self.train_data = train_data
  515. self.train_labels = train_labels
  516. self.label_encoder = label_encoder
  517. self.reference_data = reference_data
  518. self.scaler = scaler
  519. self.is_scaler_fit = is_scaler_fit
  520. self.highly_variable_genes = highly_variable_genes
  521. self.mix_included = mix_included
  522. self.experiment_name = experiment_name
  523. self.model_type = getattr(self.config.data, "model", "CNN").lower()
  524. def create_cnn_model(self, num_output_units):
  525. """
  526. Create a Convolutional Neural Network (CNN) model using TensorFlow and the provided configuration.
  527. Args:
  528. num_output_units (int): The number of output units for the final layer of the model.
  529. Returns:
  530. model (tensorflow.python.keras.Model): The created and compiled CNN model.
  531. """
  532. cnn_config = getattr(self.config.model, "cnn", {})
  533. # Create model
  534. model = tf.keras.Sequential()
  535. # Convolutional blocks
  536. for i in range(len(cnn_config.filters)):
  537. if i == 0:
  538. # First layer needs input_shape
  539. model.add(
  540. tf.keras.layers.Conv1D(
  541. filters=cnn_config.filters[i],
  542. kernel_size=cnn_config.kernel_sizes[i],
  543. strides=cnn_config.strides[i],
  544. padding=cnn_config.paddings[i],
  545. input_shape=(1, self.train_data.shape[2]),
  546. )
  547. )
  548. else:
  549. # Subsequent layers don't need input_shape
  550. model.add(
  551. tf.keras.layers.Conv1D(
  552. filters=cnn_config.filters[i],
  553. kernel_size=cnn_config.kernel_sizes[i],
  554. strides=cnn_config.strides[i],
  555. padding=cnn_config.paddings[i],
  556. )
  557. )
  558. model.add(tf.keras.layers.BatchNormalization())
  559. model.add(tf.keras.layers.ReLU())
  560. if cnn_config.pool_sizes[i] is not None:
  561. model.add(
  562. tf.keras.layers.MaxPooling1D(
  563. pool_size=cnn_config.pool_sizes[i],
  564. strides=cnn_config.pool_strides[i],
  565. padding="same",
  566. )
  567. )
  568. # Fully connected layers
  569. model.add(tf.keras.layers.Flatten())
  570. for units in cnn_config.dense_units:
  571. model.add(
  572. tf.keras.layers.Dense(units=units, activation=cnn_config.activation)
  573. )
  574. model.add(tf.keras.layers.Dropout(rate=cnn_config.dropout_rate))
  575. # Output layer based on task type
  576. task_type = getattr(self.config.model, "task_type", "classification")
  577. if task_type == "regression":
  578. model.add(tf.keras.layers.Dense(units=1, activation="linear"))
  579. default_loss = "mse"
  580. default_metrics = ["mae"]
  581. else:
  582. model.add(
  583. tf.keras.layers.Dense(units=num_output_units, activation="softmax")
  584. )
  585. default_loss = "categorical_crossentropy"
  586. # Use AUC metric object to avoid array return values in Keras 3
  587. default_metrics = [
  588. "accuracy",
  589. tf.keras.metrics.AUC(name="auc", multi_label=False),
  590. ]
  591. # Use standard Adam optimizer for all platforms (Keras 3 compatible)
  592. learning_rate = getattr(self.config.model.training, "learning_rate", 0.001)
  593. optimizer_instance = tf.keras.optimizers.Adam(learning_rate=learning_rate)
  594. # Get loss from config
  595. cnn_config = getattr(self.config.model, "cnn", {})
  596. loss = getattr(cnn_config, "loss", default_loss)
  597. # Get metrics from config - convert string "auc" to AUC object if needed
  598. eval_config = getattr(self.config, "evaluation", {})
  599. config_metrics = eval_config.get("metrics", {})
  600. training_metrics = config_metrics.get("training", {}).get(
  601. task_type, default_metrics
  602. )
  603. # Convert metric strings to metric objects to avoid array issues in Keras 3
  604. if isinstance(training_metrics, list):
  605. converted_metrics = []
  606. for m in training_metrics:
  607. if m == "auc":
  608. converted_metrics.append(
  609. tf.keras.metrics.AUC(name="auc", multi_label=False)
  610. )
  611. elif m == "precision":
  612. converted_metrics.append(
  613. tf.keras.metrics.Precision(name="precision")
  614. )
  615. elif m == "recall":
  616. converted_metrics.append(tf.keras.metrics.Recall(name="recall"))
  617. elif m == "f1_score":
  618. # F1 score needs custom implementation - skip for now
  619. pass
  620. else:
  621. converted_metrics.append(m)
  622. training_metrics = converted_metrics
  623. # Compile model
  624. model.compile(
  625. optimizer=optimizer_instance,
  626. loss=loss,
  627. metrics=training_metrics,
  628. )
  629. return model
  630. def create_mlp_model(self, num_output_units):
  631. """
  632. Create a Multilayer Perceptron (MLP) model using TensorFlow and the provided configuration.
  633. Args:
  634. num_output_units (int): The number of output units for the final layer of the model.
  635. Returns:
  636. model (tensorflow.python.keras.Model): The created and compiled MLP model.
  637. """
  638. mlp_config = getattr(self.config.model, "mlp", {})
  639. model = tf.keras.Sequential()
  640. # Input layer
  641. model.add(tf.keras.layers.InputLayer(input_shape=(self.train_data.shape[1],)))
  642. # Fully connected layers
  643. for units in mlp_config.units:
  644. model.add(
  645. tf.keras.layers.Dense(
  646. units=units, activation=mlp_config.activation_function
  647. )
  648. )
  649. model.add(tf.keras.layers.Dropout(rate=mlp_config.dropout_rate))
  650. # Output layer based on task type
  651. task_type = getattr(self.config.model, "task_type", "classification")
  652. if task_type == "regression":
  653. model.add(tf.keras.layers.Dense(units=1, activation="linear"))
  654. default_loss = "mse"
  655. default_metrics = ["mae"]
  656. else:
  657. model.add(
  658. tf.keras.layers.Dense(units=num_output_units, activation="softmax")
  659. )
  660. default_loss = "categorical_crossentropy"
  661. # Use AUC metric object to avoid array return values in Keras 3
  662. default_metrics = [
  663. "accuracy",
  664. tf.keras.metrics.AUC(name="auc", multi_label=False),
  665. ]
  666. # Use standard Adam optimizer for all platforms (Keras 3 compatible)
  667. optimizer_instance = tf.keras.optimizers.Adam(
  668. learning_rate=mlp_config.learning_rate
  669. )
  670. # Get loss from config
  671. loss = getattr(mlp_config, "loss", default_loss)
  672. # Get metrics from config - convert string "auc" to AUC object if needed
  673. eval_config = getattr(self.config, "evaluation", {})
  674. config_metrics = eval_config.get("metrics", {})
  675. training_metrics = config_metrics.get("training", {}).get(
  676. task_type, default_metrics
  677. )
  678. # Convert metric strings to metric objects to avoid array issues in Keras 3
  679. if isinstance(training_metrics, list):
  680. converted_metrics = []
  681. for m in training_metrics:
  682. if m == "auc":
  683. converted_metrics.append(
  684. tf.keras.metrics.AUC(name="auc", multi_label=False)
  685. )
  686. elif m == "precision":
  687. converted_metrics.append(
  688. tf.keras.metrics.Precision(name="precision")
  689. )
  690. elif m == "recall":
  691. converted_metrics.append(tf.keras.metrics.Recall(name="recall"))
  692. elif m == "f1_score":
  693. # F1 score needs custom implementation - skip for now
  694. pass
  695. else:
  696. converted_metrics.append(m)
  697. training_metrics = converted_metrics
  698. # Compile model
  699. model.compile(
  700. optimizer=optimizer_instance,
  701. loss=loss,
  702. metrics=training_metrics,
  703. )
  704. return model
  705. def create_logistic_regression(self):
  706. """
  707. Create a regression model using scikit-learn and the provided configuration.
  708. Returns:
  709. lr: The regression model (linear or logistic based on task type).
  710. """
  711. lr_config = getattr(
  712. self.config.model, "logistic", {}
  713. ) # Note: using "logistic" to match config
  714. task_type = getattr(self.config.model, "task_type", "classification")
  715. if task_type == "regression":
  716. from sklearn.linear_model import LinearRegression
  717. lr = LinearRegression()
  718. else:
  719. lr = LogisticRegression(
  720. penalty=getattr(lr_config, "penalty", "l2"),
  721. solver=getattr(lr_config, "solver", "lbfgs"),
  722. max_iter=getattr(lr_config, "max_iter", 1000),
  723. C=getattr(lr_config, "C", 1.0),
  724. random_state=getattr(lr_config, "random_state", 42),
  725. )
  726. return lr
  727. def create_random_forest(self):
  728. """
  729. Create a random forest model using scikit-learn and the provided configuration.
  730. Returns:
  731. rf: The random forest model (classifier or regressor based on task type).
  732. """
  733. rf_config = getattr(self.config.model, "random_forest", {})
  734. task_type = getattr(self.config.model, "task_type", "classification")
  735. if task_type == "regression":
  736. from sklearn.ensemble import RandomForestRegressor
  737. rf = RandomForestRegressor(
  738. n_estimators=rf_config.n_estimators,
  739. max_depth=rf_config.max_depth,
  740. min_samples_split=rf_config.min_samples_split,
  741. min_samples_leaf=rf_config.min_samples_leaf,
  742. max_features=rf_config.max_features,
  743. bootstrap=rf_config.bootstrap,
  744. oob_score=rf_config.oob_score,
  745. n_jobs=rf_config.n_jobs,
  746. random_state=rf_config.random_state,
  747. )
  748. else:
  749. rf = RandomForestClassifier(
  750. n_estimators=rf_config.n_estimators,
  751. criterion=rf_config.criterion,
  752. max_depth=rf_config.max_depth,
  753. min_samples_split=rf_config.min_samples_split,
  754. min_samples_leaf=rf_config.min_samples_leaf,
  755. max_features=rf_config.max_features,
  756. bootstrap=rf_config.bootstrap,
  757. oob_score=rf_config.oob_score,
  758. n_jobs=rf_config.n_jobs,
  759. random_state=rf_config.random_state,
  760. )
  761. return rf
  762. def create_xgboost_model(self):
  763. """
  764. Create an XGBoost classifier using xgboost and the provided configuration.
  765. Returns:
  766. model (xgboost.XGBClassifier): The XGBoost classifier.
  767. """
  768. xgb_config = getattr(self.config.model, "xgboost", {})
  769. # Determine task type and set corresponding XGBoost parameters
  770. task_type = getattr(self.config.model, "task_type", "classification")
  771. if task_type == "regression":
  772. xgb_objective = "reg:squarederror"
  773. eval_metric = getattr(xgb_config, "eval_metric", "rmse")
  774. else:
  775. # Classification - determine if binary or multiclass based on number of classes
  776. num_classes = (
  777. len(np.unique(self.train_labels))
  778. if hasattr(self, "train_labels")
  779. else 2
  780. )
  781. if num_classes == 2:
  782. xgb_objective = "binary:logistic"
  783. eval_metric = getattr(xgb_config, "eval_metric", "auc")
  784. else:
  785. xgb_objective = "multi:softmax"
  786. eval_metric = getattr(xgb_config, "eval_metric", "mlogloss")
  787. # Basic XGBoost parameters
  788. xgb_params = {
  789. "objective": xgb_objective,
  790. "eval_metric": eval_metric,
  791. "learning_rate": xgb_config.learning_rate,
  792. "n_estimators": xgb_config.n_estimators,
  793. "max_depth": xgb_config.max_depth,
  794. "min_child_weight": xgb_config.min_child_weight,
  795. "subsample": xgb_config.subsample,
  796. "colsample_bytree": xgb_config.colsample_bytree,
  797. "random_state": xgb_config.random_state,
  798. "tree_method": xgb_config.tree_method,
  799. "predictor": xgb_config.predictor,
  800. }
  801. # Initialize XGBoost model based on task type
  802. if task_type == "regression":
  803. model = xgb.XGBRegressor(**xgb_params)
  804. else:
  805. # Adjust parameters for multiclass classification
  806. if xgb_objective == "multi:softmax":
  807. xgb_params["num_class"] = len(np.unique(self.train_labels))
  808. model = xgb.XGBClassifier(**xgb_params)
  809. model.set_params(early_stopping_rounds=xgb_config.early_stopping_rounds)
  810. return model
  811. def build_model(self):
  812. """
  813. Build a model based on the specified type in the config.
  814. Returns:
  815. model: The created model, which could be CNN, MLP, or logistic regression.
  816. """
  817. num_output_units = (
  818. self.train_labels.shape[1] if self.model_type in ["cnn", "mlp"] else None
  819. )
  820. if self.model_type == "cnn":
  821. model = self.create_cnn_model(num_output_units)
  822. elif self.model_type == "mlp":
  823. model = self.create_mlp_model(num_output_units)
  824. elif self.model_type == "logisticregression":
  825. model = self.create_logistic_regression()
  826. elif self.model_type == "randomforest":
  827. model = self.create_random_forest()
  828. elif self.model_type == "xgboost":
  829. model = self.create_xgboost_model()
  830. else:
  831. raise ValueError("Unsupported model type provided.")
  832. return model
  833. def train_model(self, model):
  834. """
  835. Train a model and save it if it's the best one based on validation accuracy.
  836. """
  837. # Prepare directories and paths
  838. model_dir = self._prepare_directories()
  839. paths = self._define_paths(model_dir)
  840. if self.model_type in ["cnn", "mlp"]:
  841. history, model_improved = self._train_neural_network(model, paths)
  842. else:
  843. history, model_improved = self._train_sklearn_model(model, paths)
  844. return history, model, model_improved
  845. def _prepare_directories(self):
  846. """
  847. Prepare directories for saving models and related artifacts in experiment structure.
  848. Returns:
  849. experiment_dir (str): The directory where the experiment and artifacts will be saved.
  850. """
  851. self.path_manager = PathManager(self.config)
  852. # Use experiment directory instead of old model directory
  853. if self.experiment_name:
  854. experiment_dir = self.path_manager.get_experiment_dir(self.experiment_name)
  855. else:
  856. experiment_dir = self.path_manager.get_experiment_dir()
  857. os.makedirs(experiment_dir, exist_ok=True)
  858. return experiment_dir
  859. def _define_paths(self, experiment_dir):
  860. """
  861. Define paths for saving model checkpoints and related artifacts in experiment structure.
  862. Args:
  863. experiment_dir (str): The experiment directory where everything will be saved.
  864. Returns:
  865. dict: A dictionary containing paths for various artifacts.
  866. """
  867. # Create model_components subdirectory for cleaner organization
  868. components_dir = os.path.join(experiment_dir, "model_components")
  869. training_dir = os.path.join(experiment_dir, "training")
  870. os.makedirs(components_dir, exist_ok=True)
  871. os.makedirs(training_dir, exist_ok=True)
  872. paths = {
  873. "label_path": os.path.join(components_dir, "label_encoder.pkl"),
  874. "reference_path": os.path.join(components_dir, "reference_data.npy"),
  875. "scaler_path": os.path.join(components_dir, "scaler.pkl"),
  876. "is_scaler_fit_path": os.path.join(components_dir, "is_scaler_fit.pkl"),
  877. "num_features_path": os.path.join(components_dir, "num_features.pkl"),
  878. "highly_variable_genes_path": os.path.join(
  879. components_dir, "highly_variable_genes.pkl"
  880. ),
  881. "mix_included_path": os.path.join(components_dir, "mix_included.pkl"),
  882. "history_path": os.path.join(training_dir, "history.pkl"),
  883. "metadata_path": os.path.join(experiment_dir, "metadata.json"),
  884. "experiment_dir": experiment_dir,
  885. "components_dir": components_dir,
  886. }
  887. return paths
  888. def _train_neural_network(self, model, paths):
  889. """
  890. Train a neural network model (CNN or MLP) and save the best model based on validation loss.
  891. Args:
  892. model: The neural network model to train.
  893. paths (dict): A dictionary containing paths for saving artifacts.
  894. Returns:
  895. history: The training history.
  896. """
  897. custom_model_path = os.path.join(paths["experiment_dir"], "model.keras")
  898. # Try to load best validation loss from best symlink first, then current experiment
  899. from pathlib import Path
  900. # Get the correct best symlink path: base/best/config_key/model_components/best_val_loss.json
  901. config_key = self.path_manager.get_config_key()
  902. # Get base path and manually construct the correct best path
  903. base_path = Path(self.path_manager._get_project_root()) / "outputs"
  904. project_name = getattr(self.path_manager.config, "project", "fruitfly_aging")
  905. batch_correction_enabled = getattr(
  906. self.path_manager.config.data.batch_correction, "enabled", False
  907. )
  908. correction_dir = (
  909. "batch_corrected" if batch_correction_enabled else "uncorrected"
  910. )
  911. task_type = getattr(
  912. self.path_manager.config.model, "task_type", "classification"
  913. )
  914. # best_symlink_path variable removed - was unused
  915. current_path = os.path.join(paths["components_dir"], "best_val_loss.json")
  916. self.num_features = (
  917. self.train_data.shape[2]
  918. if self.model_type == "cnn"
  919. else self.train_data.shape[1]
  920. )
  921. # Load the best validation loss from all experiments
  922. best_val_loss = float("inf")
  923. # First, scan all experiments in all_runs to find the true historical best
  924. # This is more reliable than relying on potentially broken symlinks
  925. try:
  926. all_runs_path = str(
  927. base_path
  928. / project_name
  929. / "experiments"
  930. / correction_dir
  931. / task_type
  932. / "all_runs"
  933. / config_key
  934. )
  935. if os.path.exists(all_runs_path) and os.path.isdir(all_runs_path):
  936. for experiment_dir in sorted(os.listdir(all_runs_path)):
  937. if experiment_dir.startswith("experiment_"):
  938. exp_best_val_path = os.path.join(
  939. all_runs_path,
  940. experiment_dir,
  941. "model_components",
  942. "best_val_loss.json",
  943. )
  944. # Check if file exists and is not a broken symlink
  945. if os.path.exists(exp_best_val_path) and os.path.isfile(
  946. exp_best_val_path
  947. ):
  948. try:
  949. with open(exp_best_val_path) as f:
  950. exp_val_loss = json.load(f)["best_val_loss"]
  951. if exp_val_loss < best_val_loss:
  952. best_val_loss = exp_val_loss
  953. except (FileNotFoundError, json.JSONDecodeError, OSError):
  954. continue
  955. except (OSError, Exception):
  956. # Silently continue if we can't scan the all_runs directory
  957. pass
  958. # If we still don't have a best val loss, try the current path (but not broken symlinks)
  959. if best_val_loss == float("inf") and current_path:
  960. try:
  961. if os.path.exists(current_path) and os.path.isfile(current_path):
  962. with open(current_path) as f:
  963. loaded_val_loss = json.load(f)["best_val_loss"]
  964. if loaded_val_loss < best_val_loss:
  965. best_val_loss = loaded_val_loss
  966. except (FileNotFoundError, json.JSONDecodeError, OSError):
  967. pass
  968. # Previous model info will be shown in TRAINING PROGRESS header
  969. # Split data into training and validation sets
  970. # For stratification, we need 1D labels (not one-hot encoded)
  971. stratify_labels = (
  972. np.argmax(self.train_labels, axis=1)
  973. if len(self.train_labels.shape) > 1 and self.train_labels.shape[1] > 1
  974. else self.train_labels
  975. )
  976. (
  977. train_inputs_split,
  978. val_inputs_split,
  979. train_labels_split,
  980. val_labels_split,
  981. ) = train_test_split(
  982. self.train_data,
  983. self.train_labels,
  984. test_size=getattr(self.config.model.training, "validation_split", 0.2),
  985. random_state=getattr(self.config.general, "random_state", 42),
  986. stratify=stratify_labels,
  987. )
  988. # Ensure all arrays are proper numpy arrays (not masked arrays or similar)
  989. train_inputs_split = np.asarray(train_inputs_split)
  990. train_labels_split = np.asarray(train_labels_split)
  991. val_inputs_split = np.asarray(val_inputs_split)
  992. val_labels_split = np.asarray(val_labels_split)
  993. # Define callbacks for early stopping and model saving
  994. early_stopping = EarlyStopping(
  995. monitor="val_loss",
  996. patience=getattr(self.config.model.training, "early_stopping_patience", 10),
  997. verbose=0,
  998. )
  999. model_checkpoint = CustomModelCheckpoint(
  1000. custom_model_path,
  1001. current_path,
  1002. paths["label_path"],
  1003. self.label_encoder,
  1004. paths["reference_path"],
  1005. self.reference_data,
  1006. self.scaler,
  1007. paths["scaler_path"],
  1008. self.is_scaler_fit,
  1009. paths["is_scaler_fit_path"],
  1010. self.highly_variable_genes,
  1011. paths["highly_variable_genes_path"],
  1012. self.mix_included,
  1013. paths["mix_included_path"],
  1014. self.num_features,
  1015. paths["num_features_path"],
  1016. paths["metadata_path"],
  1017. path_manager=PathManager(self.config),
  1018. monitor="val_loss",
  1019. save_best_only=True,
  1020. verbose=0, # Suppress checkpoint messages
  1021. )
  1022. model_checkpoint.set_best_val_loss(best_val_loss)
  1023. # Fit the model with validation split
  1024. history = model.fit(
  1025. train_inputs_split,
  1026. train_labels_split,
  1027. epochs=getattr(self.config.model.training, "epochs", 100),
  1028. batch_size=getattr(self.config.model.training, "batch_size", 32),
  1029. validation_data=(val_inputs_split, val_labels_split),
  1030. callbacks=[early_stopping, model_checkpoint],
  1031. verbose=1, # Show progress bar per epoch
  1032. )
  1033. # Save history to models/ folder if model improved
  1034. if model_checkpoint.model_improved:
  1035. models_dir = PathManager(self.config).get_models_folder_path()
  1036. os.makedirs(models_dir, exist_ok=True)
  1037. # Save history
  1038. history_path = os.path.join(models_dir, "history.pkl")
  1039. with open(history_path, "wb") as f:
  1040. pickle.dump(history.history, f)
  1041. return history, model_checkpoint.model_improved
  1042. def _train_sklearn_model(self, model, paths):
  1043. """
  1044. Train a scikit-learn model (Logistic Regression, Random Forest, or XGBoost) and save the best model based on validation accuracy.
  1045. Args:
  1046. model: The scikit-learn model to train.
  1047. paths (dict): A dictionary containing paths for saving artifacts.
  1048. Returns:
  1049. history: The training history (if applicable).
  1050. """
  1051. train_labels = np.argmax(self.train_labels, axis=1)
  1052. # Calculate the validation split index
  1053. (
  1054. train_inputs_split,
  1055. val_inputs_split,
  1056. train_labels_split,
  1057. val_labels_split,
  1058. ) = train_test_split(
  1059. self.train_data,
  1060. train_labels,
  1061. test_size=getattr(self.config.model.training, "validation_split", 0.2),
  1062. random_state=getattr(self.config.general, "random_state", 42),
  1063. stratify=train_labels,
  1064. )
  1065. if self.model_type in ["logisticregression", "randomforest"]:
  1066. model.fit(train_inputs_split, train_labels_split)
  1067. history = None
  1068. elif self.model_type == "xgboost":
  1069. # Training with evaluation set
  1070. eval_set = [
  1071. (train_inputs_split, train_labels_split),
  1072. (val_inputs_split, val_labels_split),
  1073. ]
  1074. model.fit(
  1075. train_inputs_split,
  1076. train_labels_split,
  1077. eval_set=eval_set,
  1078. verbose=0, # Suppress XGBoost training output
  1079. )
  1080. history = model.evals_result()
  1081. else:
  1082. raise ValueError("Unsupported model type provided.")
  1083. # Evaluate the model
  1084. val_predictions = model.predict(val_inputs_split)
  1085. val_accuracy = accuracy_score(val_labels_split, val_predictions)
  1086. # Load the best validation accuracy from best symlink first, then current experiment
  1087. from pathlib import Path
  1088. # Get the correct best symlink path: base/best/config_key/model_components/best_val_accuracy.json
  1089. config_key = self.path_manager.get_config_key()
  1090. # Get base path and manually construct the correct best path
  1091. base_path = Path(self.path_manager._get_project_root()) / "outputs"
  1092. project_name = getattr(self.path_manager.config, "project", "fruitfly_aging")
  1093. batch_correction_enabled = getattr(
  1094. self.path_manager.config.data.batch_correction, "enabled", False
  1095. )
  1096. correction_dir = (
  1097. "batch_corrected" if batch_correction_enabled else "uncorrected"
  1098. )
  1099. task_type = getattr(
  1100. self.path_manager.config.model, "task_type", "classification"
  1101. )
  1102. # best_symlink_accuracy_path variable removed - was unused
  1103. current_accuracy_path = os.path.join(
  1104. paths["components_dir"], "best_val_accuracy.json"
  1105. )
  1106. # Load the best validation accuracy from file
  1107. historical_best_accuracy = 0
  1108. # First, scan all experiments in all_runs to find the true historical best
  1109. # This is more reliable than relying on potentially broken symlinks
  1110. try:
  1111. all_runs_path = str(
  1112. base_path
  1113. / project_name
  1114. / "experiments"
  1115. / correction_dir
  1116. / task_type
  1117. / "all_runs"
  1118. / config_key
  1119. )
  1120. if os.path.exists(all_runs_path) and os.path.isdir(all_runs_path):
  1121. for experiment_dir in sorted(os.listdir(all_runs_path)):
  1122. if experiment_dir.startswith("experiment_"):
  1123. exp_best_acc_path = os.path.join(
  1124. all_runs_path,
  1125. experiment_dir,
  1126. "model_components",
  1127. "best_val_accuracy.json",
  1128. )
  1129. # Check if file exists and is not a broken symlink
  1130. if os.path.exists(exp_best_acc_path) and os.path.isfile(
  1131. exp_best_acc_path
  1132. ):
  1133. try:
  1134. with open(exp_best_acc_path) as f:
  1135. exp_val_acc = json.load(f)["best_val_accuracy"]
  1136. if exp_val_acc > historical_best_accuracy:
  1137. historical_best_accuracy = exp_val_acc
  1138. except (FileNotFoundError, json.JSONDecodeError, OSError):
  1139. continue
  1140. except (OSError, Exception):
  1141. # Silently continue if we can't scan the all_runs directory
  1142. pass
  1143. # If we still don't have a best accuracy, try the current path (but not broken symlinks)
  1144. if historical_best_accuracy == 0 and current_accuracy_path:
  1145. try:
  1146. if os.path.exists(current_accuracy_path) and os.path.isfile(
  1147. current_accuracy_path
  1148. ):
  1149. with open(current_accuracy_path) as f:
  1150. loaded_val_acc = json.load(f)["best_val_accuracy"]
  1151. if loaded_val_acc > historical_best_accuracy:
  1152. historical_best_accuracy = loaded_val_acc
  1153. except (FileNotFoundError, json.JSONDecodeError, OSError):
  1154. pass
  1155. # best_val_accuracy variable removed - was unused
  1156. # Check if the current model outperforms the historical best (not just current best)
  1157. print("Validation accuracy:", val_accuracy)
  1158. print("Best validation accuracy so far:", historical_best_accuracy)
  1159. model_improved = val_accuracy > historical_best_accuracy
  1160. print("Model improved:", model_improved)
  1161. if model_improved:
  1162. # Update the record with the new best validation accuracy
  1163. with open(current_accuracy_path, "w") as f:
  1164. json.dump({"best_val_accuracy": val_accuracy}, f)
  1165. # Save selected features for reference in the future
  1166. num_features = self.train_data.shape[1]
  1167. with open(paths["num_features_path"], "wb") as f:
  1168. pickle.dump(num_features, f)
  1169. # Save the model as it's the best one so far using pickle
  1170. model_path = os.path.join(paths["experiment_dir"], "model.pkl")
  1171. with open(model_path, "wb") as f:
  1172. pickle.dump(model, f)
  1173. # XGBoost history is saved to models/ folder in pipeline_manager after training
  1174. # Save the label encoder
  1175. with open(paths["label_path"], "wb") as label_file:
  1176. pickle.dump(
  1177. self.label_encoder, label_file
  1178. ) # save the label_encoder when the model improves
  1179. # Save reference data
  1180. np.save(paths["reference_path"], self.reference_data)
  1181. # Save scaler
  1182. with open(paths["scaler_path"], "wb") as scaler_file:
  1183. pickle.dump(self.scaler, scaler_file)
  1184. # Save is_scaler_fit
  1185. with open(paths["is_scaler_fit_path"], "wb") as is_scaler_fit_file:
  1186. pickle.dump(self.is_scaler_fit, is_scaler_fit_file)
  1187. with open(paths["mix_included_path"], "wb") as mix_included_file:
  1188. pickle.dump(self.mix_included, mix_included_file)
  1189. # Save highly variable genes
  1190. with open(
  1191. paths["highly_variable_genes_path"], "wb"
  1192. ) as highly_variable_genes_file:
  1193. pickle.dump(self.highly_variable_genes, highly_variable_genes_file)
  1194. print("New best model saved with validation accuracy:", val_accuracy)
  1195. # Save history to models/ folder for XGBoost (if it exists)
  1196. if self.model_type == "xgboost" and history is not None:
  1197. models_dir = PathManager(self.config).get_models_folder_path()
  1198. os.makedirs(models_dir, exist_ok=True)
  1199. history_path = os.path.join(models_dir, "history.pkl")
  1200. with open(history_path, "wb") as f:
  1201. pickle.dump(history, f)
  1202. return history, model_improved
  1203. def run(self):
  1204. """
  1205. Builds and trains the model based on the configuration settings.
  1206. This method constructs the model using the specified architecture (e.g., CNN, MLP),
  1207. and then trains it using the training data.
  1208. Returns:
  1209. - tuple: A tuple containing the trained model and the training history.
  1210. """
  1211. # Build the model using the provided configuration
  1212. model = self.build_model()
  1213. # Train the model using the provided training data and additional components
  1214. history, model, model_improved = self.train_model(model)
  1215. return model, history, model_improved

model.py at commit c5c6e5e, under CC-BY-NC-ND-4.0 · at the source

Overview

Authors: Nikolai Tennant1, Ananya Pavuluri2, Gunjan Singh3,4, Kaitlyn Cortez3, Kate O’Connor-Giles4,5, Erica Larschan2,3, Ritambhara Singh1,2,6
  1. Data Science Institute, Brown University,Providence, RI USA
  2. Center for Computational Molecular Biology, Brown University,Providence, RI USA
  3. Department of Molecular Biology, Cell Biology, and Biochemistry, Brown University,Providence, RI USA
  4. Department of Neuroscience, Brown University,Providence, RI USA
  5. Carney Institute for Brain Science, Brown University,Providence, RI USA
  6. Department of Computer Science, Brown University,Providence, RI USA
Institutions: Brown University (United States)
Journal: Scientific reports, volume 16, issue 1, article 17434
Dates: received 16 November 2024; accepted 8 April 2026; published online 14 April 2026
Type: Research article · Language: English
License: CC BY-NC-ND
Identifiers: DOI 10.1038/s41598-026-48613-0 · PMID 41981187 · PMCID PMC13237136 · OpenAlex W7154252694
Open access: gold, a free copy (OpenAlex)
Status: code verified
Categories: genetics / omics (modality), drosophila (organism), cellular / molecular (subfield)
Methods: Statistics, Machine learning, Evoked potentials
Keywords: Drosophila melanogaster, Transcriptomics, Deep learning, Single-cell, Aging clock, Dosage compensation, Machine learning, Ageing, Gene expression
MeSH: Aging*, Dosage Compensation, Genetic*, Drosophila melanogaster*, RNA, Long Noncoding*, Animals, Deep Learning, Female, Head, Male, Sex Factors, Single-Cell Gene Expression Analysis, Time Factors (* major topic)
Topic: Single-cell and spatial transcriptomics (Molecular Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Citations: cited by 1 paper (Europe PMC); 68 references in the paper

Abstract

The abstract is not reproduced here: the paper's license (CC BY-NC-ND) does not allow it. Read it in the paper, at the publisher or on Europe PMC.

Repository

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

rsinghlab/TimeFlies

License: CC-BY-NC-ND-4.0
State: the link answers, verified on 29 September 2026
Evidence: files inventoried
Commit: c5c6e5e35621487514bc11039abb456313b2ba80, 2 June 2026
Languages: Python (70)
Size: 89 files, 70 scripts
Software Heritage: not archived
Found in: “Code availability”
Holds: README, license file, environment (pyproject.toml, uv.lock, src/timeflies/data/setup.py, src/timeflies/cli/commands/setup.py), tests, continuous integration
Not found: CITATION.cff, documentation
Tools: NumPy (23 files), pandas (15 files), anndata (13 files), scikit-learn (13 files), TensorFlow (8 files), Scanpy (7 files), Matplotlib (5 files), Keras (4 files), SciPy (4 files), seaborn (3 files), SHAP (2 files), XGBoost (2 files), PyTorch (1 file)
Availability: 1 check, the latest on 29 September 2026: the link answers
  • 29 September 2026: the link answers
72 files

Code availability statement

The paper has a code availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

Read it in the paper: doi.org/10.1038/s41598-026-48613-0.

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;
  • 70 scripts, each with its path and the digest of its content;
  • 4 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 availability statement

The paper has a data availability statement. Its license (CC BY-NC-ND) does not allow reproducing it here; in short, from what the harvester recognized in it:

  • no repository, dataset or request procedure was recognized in it

Read it in the paper: doi.org/10.1038/s41598-026-48613-0.

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, 29 September 2026: the first record

Recorded: type, language, journal, volume, issue, pages, dates, 7 authors, 9 keywords, 12 MeSH terms, 3 funders, 67 references.

Cite

This paper

Tennant, N., Pavuluri, A., Singh, G., Cortez, K., O’Connor-Giles, K., Larschan, E., & Singh, R. (2026). An snRNA-seq aging clock for the fruit fly head sheds light on sex-biased aging. Scientific reports, 16(1), 17434. https://doi.org/10.1038/s41598-026-48613-0

BibTeX

@article{tennant2026snrna,
author = {Tennant, Nikolai and Pavuluri, Ananya and Singh, Gunjan and Cortez, Kaitlyn and O’Connor-Giles, Kate and Larschan, Erica and Singh, Ritambhara},
title = {{An snRNA-seq aging clock for the fruit fly head sheds light on sex-biased aging}},
journal = {Scientific reports},
year = {2026},
month = apr,
volume = {16},
number = {1},
pages = {17434},
publisher = {Nature Publishing Group},
issn = {2045-2322},
doi = {10.1038/s41598-026-48613-0},
url = {https://doi.org/10.1038/s41598-026-48613-0},
pmid = {41981187},
pmcid = {PMC13237136}
}

RIS

TY - JOUR
AU - Tennant, Nikolai
AU - Pavuluri, Ananya
AU - Singh, Gunjan
AU - Cortez, Kaitlyn
AU - O’Connor-Giles, Kate
AU - Larschan, Erica
AU - Singh, Ritambhara
TI - An snRNA-seq aging clock for the fruit fly head sheds light on sex-biased aging
T2 - Scientific reports
J2 - Sci Rep
PY - 2026
DA - 2026/04/14
VL - 16
IS - 1
SP - 17434
SN - 2045-2322
PB - Nature Publishing Group
DO - 10.1038/s41598-026-48613-0
UR - https://doi.org/10.1038/s41598-026-48613-0
LA - en
ER -

CSL-JSON

{
"id": "10.1038/s41598-026-48613-0",
"type": "article-journal",
"title": "An snRNA-seq aging clock for the fruit fly head sheds light on sex-biased aging",
"container-title": "Scientific reports",
"author": [
{
"family": "Tennant",
"given": "Nikolai"
},
{
"family": "Pavuluri",
"given": "Ananya"
},
{
"family": "Singh",
"given": "Gunjan"
},
{
"family": "Cortez",
"given": "Kaitlyn"
},
{
"family": "O’Connor-Giles",
"given": "Kate"
},
{
"family": "Larschan",
"given": "Erica"
},
{
"family": "Singh",
"given": "Ritambhara"
}
],
"container-title-short": "Sci Rep",
"volume": "16",
"issue": "1",
"page": "17434",
"DOI": "10.1038/s41598-026-48613-0",
"PMID": "41981187",
"PMCID": "PMC13237136",
"ISSN": "2045-2322",
"publisher": "Nature Publishing Group",
"URL": "https://doi.org/10.1038/s41598-026-48613-0",
"language": "en",
"issued": {
"date-parts": [
[
2026,
4,
14
]
]
}
}

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.1016/j.isci.2026.116439 [code]
Decoding the role of transcriptomic clocks in the human prefrontal cortex.
Journal: iScience
In common: Keras, TensorFlow, seaborn, 5 other tools, genetics / omics, 8 references
[2] 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, anndata, Scanpy, 8 other tools, genetics / omics, 3 references
[3] doi:10.1038/s42003-026-10462-y [code]
SpaDC enables sequence-based integrative analysis and regulatory inference of spatial chromatin accessibility data.
Journal: Communications biology
In common: SHAP, XGBoost, anndata, 7 other tools, genetics / omics, 3 references
[4] doi:10.1038/s41514-026-00391-9 [code]
Region-specific transcriptional signatures of brain aging in the absence of neuropathology at the single-cell level.
Journal: npj aging
In common: anndata, Scanpy, PyTorch, 6 other tools, genetics / omics, cellular / molecular, 4 references
[5] doi:10.1038/s44320-026-00208-7 [code]
Interpretable deep generative ensemble learning for single-cell omics with Hydra.
Journal: Molecular systems biology
In common: Keras, anndata, Scanpy, 8 other tools, cellular / molecular, 1 reference
[6] doi:10.1016/j.xgen.2026.101217 [code]
ProtoCloud: A prototypical self-explaining model for single-cell analysis.
Journal: Cell genomics
In common: anndata, Scanpy, PyTorch, 6 other tools, genetics / omics, cellular / molecular, 3 references
[7] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: SHAP, XGBoost, Keras, 8 other tools, cellular / molecular
[8] doi:10.1523/eneuro.0362-25.2026 [code]
Similarities between &lt;i&gt;Ciona&lt;/i&gt; Dorsal Motor Ganglion and Vertebrate Cerebellum: Did a Chordate Ancestor Already Show D/V Subdivision within a Hindbrain Precursor?
Journal: eNeuro
In common: SHAP, anndata, Scanpy, 7 other tools, 2 references
[9] doi:10.1039/d6ra03343a [code]
A benchmark dataset and interpretable deep learning framework for drug-induced developmental neurotoxicity prediction.
Journal: RSC advances
In common: SHAP, Keras, TensorFlow, 6 other tools, 2 references
[10] 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, XGBoost, Keras, 8 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.