Source code for rtnn.dataset_odepth

# Copyright 2026 IPSL / CNRS / Sorbonne University
# Authors: Kazem Ardaneh
#
# This work is licensed under the Creative Commons
# Attribution-NonCommercial-ShareAlike 4.0 International License.
# To view a copy of this license, visit
# http://creativecommons.org/licenses/by-nc-sa/4.0/

import torch
from torch.utils.data import Dataset
import numpy as np
import xarray as xr
from typing import Dict, List, Tuple, Any
import random


[docs] class RRTMGPGASDataPreprocessor(Dataset): """ Dataset class for gas optics emulator. Inputs (7 features): tlay, play, h2o, o3, co2, n2o, ch4 Outputs (predictand): tau_sw_abs, tau_sw_ray, ssa_sw, etc. For each experiment (expt), returns data in RNN format (B, C, T): - B: n_sites_per_batch * n_gpt (batch dimension) - C: n_features = 7 (tlay, play, h2o, o3, co2, n2o, ch4) - T: n_layer = 60 (sequence length) Site selection: - Training: Random start index per batch per experiment with periodic boundary - Validation: Sequential all sites without wrap Parameters ---------- logger : Any Logger instance path : str Path to NetCDF file predictand : str Target variable: 'tau_sw_abs', 'tau_sw_ray', 'ssa_sw', 'tau_sw' training : bool If True, enable data augmentation (random sampling of sites) norm_mapping : Dict, optional Normalization statistics for each variable normalization_type : Dict, optional Normalization method for each variable debug : bool, optional Enable debug logging n_sites_per_batch : int, optional Number of sites per spatial batch (default: 128) sbatch : int, optional Number of spatial batches for training (default: 64) """
[docs] def __init__( self, logger: Any, path: str, predictand: str = "tau_sw_ray", training: bool = True, norm_mapping: Dict = {}, normalization_type: Dict = {}, debug: bool = False, n_sites_per_batch: int = 128, ) -> None: super().__init__() self.logger = logger self.path = path self.predictand = predictand self.training = training self.norm_mapping = norm_mapping self.normalization_type = normalization_type self.debug = debug self.n_sites_per_batch = n_sites_per_batch # Supported predictands self.supported_predictands = [ "tau_lw", "planck_frac", "tau_sw", "tau_sw_abs", "tau_sw_ray", "ssa_sw", ] if self.predictand not in self.supported_predictands: raise ValueError( f"predictand must be one of {self.supported_predictands}, got {self.predictand}" ) if self.debug: self.logger.info(f"Loading dataset: {path}") self.ds = xr.open_dataset(path) # Get dimensions self.n_expt = self.ds.sizes["expt"] # Usually 1 self.n_site = self.ds.sizes["site"] # 32768 self.n_layer = self.ds.sizes["layer"] # 60 self.n_level = self.ds.sizes["level"] # 61 self.n_gpt = self.ds.sizes["gpt"] # 224 for shortwave self.n_feature = self.ds.sizes["feature"] # 7 self.sbatch = self.n_site // self.n_sites_per_batch # Number of spatial batches # Total number of experiments self.n_experiments = self.n_expt # Determine number of spatial batches based on mode if self.training: self.n_batches = self.sbatch self.n_sites_used = self.n_sites_per_batch * self.sbatch self.last_expt_idx = -1 self.current_start_indices = None if self.debug: self.logger.info( f"Training: {self.n_batches} batches, {self.n_sites_used} total sites used (with periodicity)" ) else: self.n_batches = max(1, self.n_site // self.n_sites_per_batch) if self.n_site % self.n_sites_per_batch != 0: self.n_batches += 1 self.n_sites_used = self.n_site if self.debug: self.logger.info( f"Validation: {self.n_batches} batches, {self.n_sites_used} total sites" ) # Data dimensions for RNN format # Input features: tlay, play, h2o, o3, co2, n2o, ch4 (7 features) self.n_features = 7 # Outputs: depends on predictand (1 output per g-point) self.n_outputs = 1 # Each predictand is a single value per g-point # Pre-load data references self.rrtmgp_input = self.ds["rrtmgp_sw_input"] # Load predictand self._load_predictand() # Feature names for normalization (7 features) self.feature_names = ["tlay", "play", "h2o", "o3", "co2", "n2o", "ch4"] self.output_names = [self.predictand] self._logger_info()
def _load_predictand(self): """Load the target variable based on predictand.""" if self.predictand in ["tau_lw", "planck_frac"]: # Longwave - not implemented in this dataset raise NotImplementedError( f"Longwave predictand '{self.predictand}' not implemented" ) else: # Shortwave if self.predictand == "tau_sw_ray": # tau_sw_ray = tau * ssa tau = self.ds["tau_sw_gas"].values ssa = self.ds["ssa_sw_gas"].values self.y = tau * ssa del tau, ssa elif self.predictand == "tau_sw_abs": # tau_sw_abs = tau - tau * ssa tau = self.ds["tau_sw_gas"].values ssa = self.ds["ssa_sw_gas"].values tau_sw_ray = tau * ssa self.y = tau - tau_sw_ray del tau, ssa, tau_sw_ray else: # Direct variable: 'tau_sw', 'ssa_sw' self.y = self.ds[self.predictand].values def _get_random_start_indices(self) -> List[int]: """Generate random start indices for each batch.""" return [random.randint(0, self.n_site - 1) for _ in range(self.n_batches)] def _get_site_indices_for_batch(self, start_idx: int, batch_idx: int) -> List[int]: """Get site indices for a specific batch.""" if self.training: batch_start = start_idx batch_end = start_idx + self.n_sites_per_batch site_indices = [i % self.n_site for i in range(batch_start, batch_end)] return site_indices else: batch_start = batch_idx * self.n_sites_per_batch batch_end = min(batch_start + self.n_sites_per_batch, self.n_site) site_indices = list(range(batch_start, batch_end)) return site_indices def _logger_info(self): """Log dataset information.""" self.logger.info("=" * 70) self.logger.info("RRTMGPGAS DataPreprocessor (NN-RRTMGPGAS)") self.logger.info(f"File: {self.path.split('/')[-1]}") self.logger.info(f"Predictand: {self.predictand}") self.logger.info(f"Training mode: {self.training}") self.logger.info(f"Spatial batches: {self.n_batches}") self.logger.info(f"Sites per batch: {self.n_sites_per_batch}") self.logger.info( f"Total sites used: {self.n_sites_used} (out of {self.n_site})" ) self.logger.info(f"Total experiments (expt): {self.n_experiments}") self.logger.info( f"Dimensions: expt={self.n_expt}, site={self.n_site}, layer={self.n_layer}, gpt={self.n_gpt}" ) self.logger.info( f"Features: {self.n_features} (tlay, play, h2o, o3, co2, n2o, ch4)" ) self.logger.info(f"Outputs: {self.n_outputs} ({self.predictand})") self.logger.info( f"RNN format: (B={self.n_sites_per_batch}*{self.n_gpt}, C={self.n_features}, T={self.n_layer})" ) if self.training: self.logger.info( "Site selection: Periodic boundary with random start index per batch per experiment" ) else: self.logger.info("Site selection: Sequential all sites (no wrap)") self.logger.info("=" * 70)
[docs] def normalize(self, data: np.ndarray, var_name: str) -> np.ndarray: """Normalize data using stored statistics.""" if not self.norm_mapping or var_name not in self.norm_mapping: return data norm_type = self.normalization_type.get(var_name, "minmax") stats = self.norm_mapping[var_name] if norm_type == "minmax": vmin, vmax = stats["vmin"], stats["vmax"] return (data - vmin) / (vmax - vmin + 1e-8) elif norm_type == "standard": mean, std = stats["vmean"], stats["vstd"] return (data - mean) / (std + 1e-8) elif norm_type == "robust": median, iqr = stats["median"], stats["iqr"] return (data - median) / (iqr + 1e-8) elif norm_type == "log1p_standard": data_log = np.log1p(np.clip(data, a_min=0, a_max=None)) mean, std = stats["log_mean"], stats["log_std"] return (data_log - mean) / (std + 1e-8) elif norm_type == "log1p_minmax": data_log = np.log1p(np.clip(data, a_min=0, a_max=None)) vmin, vmax = stats["log_min"], stats["log_max"] return (data_log - vmin) / (vmax - vmin + 1e-8) elif norm_type == "log1p_robust": data_log = np.log1p(np.clip(data, a_min=0, a_max=None)) median, iqr = stats["log_median"], stats["log_iqr"] return (data_log - median) / (iqr + 1e-8) elif norm_type == "sqrt_standard": data_sqrt = np.sqrt(np.clip(data, a_min=0, a_max=None)) mean, std = stats["sqrt_mean"], stats["sqrt_std"] return (data_sqrt - mean) / (std + 1e-8) elif norm_type == "sqrt_minmax": data_sqrt = np.sqrt(np.clip(data, a_min=0, a_max=None)) vmin, vmax = stats["sqrt_min"], stats["sqrt_max"] return (data_sqrt - vmin) / (vmax - vmin + 1e-8) elif norm_type == "sqrt_robust": data_sqrt = np.sqrt(np.clip(data, a_min=0, a_max=None)) median, iqr = stats["sqrt_median"], stats["sqrt_iqr"] return (data_sqrt - median) / (iqr + 1e-8) else: return data
def __len__(self) -> int: """Return number of samples (experiments * spatial batches).""" return self.n_experiments * self.n_batches def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]: """ Get a specific spatial batch for a specific experiment in RNN format (B, C, T). Returns ------- Tuple[torch.Tensor, torch.Tensor] - features: (n_sites_per_batch, n_features, n_layer) - targets: (n_sites_per_batch, n_gpt, n_layer) """ batch_idx = idx % self.n_batches expt_idx = idx // self.n_batches # Determine site indices if self.training: if self.last_expt_idx != expt_idx: self.current_start_indices = self._get_random_start_indices() self.last_expt_idx = expt_idx if self.debug: self.logger.info( f"New start indices for expt {expt_idx}: first 5 = {self.current_start_indices[:5]}" ) start_idx = self.current_start_indices[batch_idx] site_indices = self._get_site_indices_for_batch(start_idx, batch_idx) else: site_indices = self._get_site_indices_for_batch(0, batch_idx) n_sites = len(site_indices) if self.debug: self.logger.info(f"\nExperiment {expt_idx}, Batch {batch_idx}:") self.logger.info(f" Number of sites: {n_sites}") self.logger.info( f" First 5 sites: {site_indices[:5] if n_sites > 0 else []}" ) # Extract features: (n_sites, n_layer, n_features) # rrtmgp_sw_input has shape (expt, site, layer, feature) features = self.rrtmgp_input[expt_idx, site_indices, :, :].values # Extract targets: (n_sites, n_layer, n_gpt) targets = self.y[expt_idx, site_indices, :, :] # Normalize features: apply to each feature for i, var_name in enumerate(self.feature_names): features[..., i] = self.normalize(features[..., i], var_name) # Normalize targets: apply to each g-point targets = self.normalize(targets, self.predictand) # Convert to RNN format (B, C, T) # features: (n_sites, n_layer, n_features) -> (n_sites, n_features, n_layer) features = np.transpose(features, (0, 2, 1)) # targets: (n_sites, n_layer, n_gpt) -> (n_sites, n_gpt, n_layer) targets = np.transpose(targets, (0, 2, 1)) # Convert to tensors features_tensor = torch.tensor(features, dtype=torch.float32) targets_tensor = torch.tensor(targets, dtype=torch.float32) if self.debug: self.logger.info( f" Features shape (B, C, T): {features_tensor.shape}" ) # (n_sites, 7, 60) self.logger.info( f" Targets shape (B, C, T): {targets_tensor.shape}" ) # (n_sites, 224, 60) return features_tensor, targets_tensor