# 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