From 21b34394c269933c54d603fe6456c76938df2d0e Mon Sep 17 00:00:00 2001 From: Zhuowei Gu Date: Wed, 3 Dec 2025 11:38:15 -0500 Subject: [PATCH 1/2] added drop scalar variables option --- config/training_config.py | 21 +- data/data_loader_individual.py | 63 +++- models/cnp_combined_model.py | 52 +-- train_cnp_model.py | 13 + training/trainer.py | 622 +++++++++++++++++++-------------- 5 files changed, 461 insertions(+), 310 deletions(-) diff --git a/config/training_config.py b/config/training_config.py index d297872..197441a 100644 --- a/config/training_config.py +++ b/config/training_config.py @@ -644,6 +644,7 @@ def get_cnp_combined_config( use_tva4km: bool = False, max_files: Optional[int] = None, include_water: bool = False, + include_scalar: bool = True, variable_list_path: Optional[str] = None, model_config_path: Optional[str] = None ) -> TrainingConfigManager: @@ -692,7 +693,7 @@ def get_cnp_combined_config( # Fallback to defaults if none provided via CNP_IO if not data_paths: if use_trendy1: - data_paths.append("/global/cfs/cdirs/m4814/daweigao/14_Code/all_dataset_1_degree") + data_paths.append("/gpfs/wolf2/cades/cli185/proj-shared/guzhuowei0407/Dataset_test/TVA/old_TVA_enhanced_dataset") if use_trendy05: data_paths.append("/mnt/proj-shared/AI4BGC_7xw/TrainingData/Trendy_05_data_CNP") if use_tva4km: @@ -704,7 +705,7 @@ def get_cnp_combined_config( if env_pat: dataset_file_patterns[env_tva] = env_pat if file_pattern is None: - file_pattern = "enhanced_1_training_data_batch_*.pkl" + file_pattern = "enhanced_monthly_training_data_batch_*.pkl" config.update_data_config( data_paths=data_paths, @@ -803,26 +804,30 @@ def get_cnp_combined_config( time_series_columns=time_series_columns, static_columns=surface_properties, pft_param_columns=pft_parameters, - x_list_scalar_columns=scalar_variables, x_list_columns_1d=pft_1d_variables, x_list_columns_2d=variables_2d_soil, longitudes_to_drop=longitudes_to_drop ) + if include_scalar: + data_config_kwargs['x_list_scalar_columns'] = scalar_variables if include_water: data_config_kwargs['x_list_water_columns'] = water_variables config.update_data_config(**data_config_kwargs) # Outputs - output_scalar = ['Y_' + v for v in scalar_variables] output_1d_pft = ['Y_' + v for v in pft_1d_variables] output_2d = ['Y_' + v for v in variables_2d_soil] - - config.update_data_config( - y_list_scalar_columns=output_scalar, + + output_config_kwargs = dict( y_list_columns_1d=output_1d_pft, y_list_columns_2d=output_2d ) + if include_scalar: + output_scalar = ['Y_' + v for v in scalar_variables] + output_config_kwargs['y_list_scalar_columns'] = output_scalar + + config.update_data_config(**output_config_kwargs) # Model configuration for CNP architecture (defaults) config.update_model_config( @@ -842,7 +847,7 @@ def get_cnp_combined_config( water_fc_size=64 if include_water else 0, # Reduced from 128 # FC for scalar variables (4 variables) - separate from surface properties - scalar_fc_size=64, # Reduced from 128 + scalar_fc_size=64 if include_scalar else 0, # Reduced from 128 # FC for 1D PFT variables (14 variables) - separate from surface properties pft_1d_fc_size=64, # Reduced from 128 diff --git a/data/data_loader_individual.py b/data/data_loader_individual.py index d9cd8f6..5bde96d 100644 --- a/data/data_loader_individual.py +++ b/data/data_loader_individual.py @@ -384,9 +384,15 @@ def normalize_data(self) -> Dict[str, Any]: # Static (keep group normalization for now) static_data, static_scaler = self._normalize_static(self.data_config.static_columns) - # Scalar - Use group normalization - scalar_data, scalar_scaler = self._normalize_scalar() - y_scalar_data, y_scalar_scaler = self._normalize_y_scalar() + # Scalar - Use group normalization (if present) + scalar_data = None + scalar_scaler = None + y_scalar_data = None + y_scalar_scaler = None + if hasattr(self.data_config, 'x_list_scalar_columns') and self.data_config.x_list_scalar_columns: + scalar_data, scalar_scaler = self._normalize_scalar() + if hasattr(self.data_config, 'y_list_scalar_columns') and self.data_config.y_list_scalar_columns: + y_scalar_data, y_scalar_scaler = self._normalize_y_scalar() # 1D PFT - Use group normalization pft_1d_data, pft_1d_scaler = self._normalize_list_1d(self.data_config.x_list_columns_1d) @@ -421,13 +427,13 @@ def normalize_data(self) -> Dict[str, Any]: self.scalers = { 'time_series': time_series_scaler, 'static': static_scaler, - 'scalar': scalar_scaler, - 'y_scalar': y_scalar_scaler, 'pft_1d': pft_1d_scaler, 'y_pft_1d': y_pft_1d_scaler, 'variables_2d_soil': variables_2d_soil_scaler, 'y_soil_2d': y_soil_2d_scaler, 'pft_param': pft_param_scaler, + 'scalar': scalar_scaler if scalar_scaler is not None else None, + 'y_scalar': y_scalar_scaler if y_scalar_scaler is not None else None, 'water': water_scaler if 'water_scaler' in locals() else None, 'y_water': y_water_scaler if 'y_water_scaler' in locals() else None, } @@ -436,16 +442,18 @@ def normalize_data(self) -> Dict[str, Any]: 'time_series_data': time_series_data, 'static_data': static_data, 'pft_param_data': pft_param_data, - 'scalar_data': scalar_data, 'variables_1d_pft': pft_1d_data, 'variables_2d_soil': variables_2d_soil, - 'y_scalar': y_scalar_data, 'y_pft_1d': y_pft_1d_data, 'y_soil_2d': y_soil_2d, 'water': water_tensor, 'y_water': y_water_tensor, 'scalers': self.scalers } + if scalar_data is not None: + ret['scalar_data'] = scalar_data + if y_scalar_data is not None: + ret['y_scalar'] = y_scalar_data # Add per-sample PFT mask derived from raw PCT_NAT_PFT_1..16 (1 where >0, else 0) try: pct_cols = [f'PCT_NAT_PFT_{i}' for i in range(1, 17)] @@ -488,9 +496,15 @@ def normalize_data_individual(self, transform_only: bool = False) -> Dict[str, A # Static (keep group normalization for now) static_data, static_scaler = self._normalize_static(self.data_config.static_columns) - # Scalar - Use individual normalization - scalar_data, scalar_scaler = self._normalize_scalar_individual(transform_only) - y_scalar_data, y_scalar_scaler = self._normalize_y_scalar_individual(transform_only) + # Scalar - Use individual normalization (if present) + scalar_data = None + scalar_scaler = None + y_scalar_data = None + y_scalar_scaler = None + if hasattr(self.data_config, 'x_list_scalar_columns') and self.data_config.x_list_scalar_columns: + scalar_data, scalar_scaler = self._normalize_scalar_individual(transform_only) + if hasattr(self.data_config, 'y_list_scalar_columns') and self.data_config.y_list_scalar_columns: + y_scalar_data, y_scalar_scaler = self._normalize_y_scalar_individual(transform_only) # 1D PFT - Use individual normalization pft_1d_data, pft_1d_scaler = self._normalize_list_1d_individual(self.data_config.x_list_columns_1d, transform_only) @@ -519,7 +533,8 @@ def normalize_data_individual(self, transform_only: bool = False) -> Dict[str, A assert pft_param_data.shape[1] == len(self.data_config.pft_param_columns), 'Mismatch in PFT param feature count!' # Only assert Y variables if they were normalized (training mode) if y_scalar_data is not None: - assert y_scalar_data.shape[1] == len(self.data_config.y_list_scalar_columns), 'Mismatch in y_scalar feature count!' + if y_scalar_data is not None: + assert y_scalar_data.shape[1] == len(self.data_config.y_list_scalar_columns), 'Mismatch in y_scalar feature count!' if y_pft_1d_data is not None: assert y_pft_1d_data.shape[1] == len(self.data_config.y_list_columns_1d), 'Mismatch in y_pft_1d variable count!' if y_soil_2d is not None: @@ -567,13 +582,14 @@ def normalize_data_individual(self, transform_only: bool = False) -> Dict[str, A 'time_series_data': time_series_data, 'static_data': static_data, 'pft_param_data': pft_param_data, - 'scalar_data': scalar_data, 'variables_1d_pft': pft_1d_data, 'variables_2d_soil': variables_2d_soil, 'water': water_tensor, 'y_water': y_water_tensor, 'scalers': self.scalers } + if scalar_data is not None: + ret['scalar_data'] = scalar_data # Only add Y variables if they were normalized (training mode) or exist (inference mode) if y_scalar_data is not None: @@ -1574,11 +1590,12 @@ def split_data(self, normalized_data: Dict[str, Any]) -> Dict[str, Any]: train_data['pft_param'] = train_pft_param test_data['pft_param'] = test_pft_param - # Split scalar data (input) - train_list_scalar = normalized_data['scalar_data'][:train_size] - test_list_scalar = normalized_data['scalar_data'][train_size:] - train_data['scalar'] = train_list_scalar - test_data['scalar'] = test_list_scalar + # Split scalar data (input) - skip if not present + if 'scalar_data' in normalized_data and normalized_data['scalar_data'] is not None: + train_list_scalar = normalized_data['scalar_data'][:train_size] + test_list_scalar = normalized_data['scalar_data'][train_size:] + train_data['scalar'] = train_list_scalar + test_data['scalar'] = test_list_scalar # Split y_scalar (target) - skip if not present (inference mode) if 'y_scalar' in normalized_data and normalized_data['y_scalar'] is not None: @@ -1602,6 +1619,18 @@ def split_data(self, normalized_data: Dict[str, Any]) -> Dict[str, Any]: y_soil_2d = normalized_data['y_soil_2d'] train_data['y_soil_2d'] = y_soil_2d[:train_size] test_data['y_soil_2d'] = y_soil_2d[train_size:] + + # Split scalar (input) - skip if not present + if 'scalar_data' in normalized_data and normalized_data['scalar_data'] is not None: + scalar = normalized_data['scalar_data'] + train_data['scalar'] = scalar[:train_size] + test_data['scalar'] = scalar[train_size:] + + # Split y_scalar (target) - skip if not present + if 'y_scalar' in normalized_data and normalized_data['y_scalar'] is not None: + y_scalar = normalized_data['y_scalar'] + train_data['y_scalar'] = y_scalar[:train_size] + test_data['y_scalar'] = y_scalar[train_size:] # Split variables_2d_soil (input) variables_2d_soil = normalized_data['variables_2d_soil'] diff --git a/models/cnp_combined_model.py b/models/cnp_combined_model.py index e448824..5134bc6 100644 --- a/models/cnp_combined_model.py +++ b/models/cnp_combined_model.py @@ -43,7 +43,7 @@ class CNPCombinedModel(nn.Module): """ def __init__(self, model_config: ModelConfig, data_info: Dict[str, Any], - include_water: bool = True, use_learnable_loss_weights: bool = False): + include_water: bool = True, include_scalar: bool = True, use_learnable_loss_weights: bool = False): """ Initialize the CNP combined model. """ @@ -52,6 +52,7 @@ def __init__(self, model_config: ModelConfig, data_info: Dict[str, Any], self.model_config = model_config self.data_info = data_info self.include_water = include_water + self.include_scalar = include_scalar self.use_learnable_loss_weights = use_learnable_loss_weights self.token_dim = self.model_config.token_dim # <-- Fix: set token_dim before feature fusion # Centralized dropout probability (allows disabling for strict determinism) @@ -77,7 +78,10 @@ def __init__(self, model_config: ModelConfig, data_info: Dict[str, Any], # Learnable log_sigma parameters for loss weighting (optional) if self.use_learnable_loss_weights: - self.log_sigma_scalar = nn.Parameter(torch.zeros(1)) + if self.include_scalar: + self.log_sigma_scalar = nn.Parameter(torch.zeros(1)) + else: + self.log_sigma_scalar = None self.log_sigma_soil_2d = nn.Parameter(torch.zeros(1)) if self.include_water: self.log_sigma_water = nn.Parameter(torch.zeros(1)) @@ -122,8 +126,11 @@ def _calculate_input_dimensions(self): else: self.water_input_size = 0 - # Scalar variables input size (4 variables) - self.scalar_input_size = len(self.data_info.get('x_list_scalar_columns', [])) + # Scalar variables input size (4 variables, optional) + if self.include_scalar: + self.scalar_input_size = len(self.data_info.get('x_list_scalar_columns', [])) + else: + self.scalar_input_size = 0 # print(f"[DEBUG] scalar_input_size at model init: {self.scalar_input_size}") # 2D input size @@ -190,8 +197,8 @@ def _build_water_encoder(self): self.fc_water = None def _build_scalar_encoder(self): - """Build encoder for scalar variables (4 variables).""" - if self.scalar_input_size > 0: + """Build encoder for scalar variables (4 variables, optional).""" + if self.include_scalar and self.scalar_input_size > 0: self.fc_scalar = nn.Sequential( nn.Linear(self.scalar_input_size, 32), nn.ReLU(), @@ -396,14 +403,17 @@ def _build_output_heads(self): ) else: self.water_head = None - # Scalar output head (6 variables) - self.scalar_head = nn.Sequential( - nn.Linear(self.token_dim, 64), - nn.BatchNorm1d(64), # Add BatchNorm - nn.ReLU(), - nn.Dropout(self.dropout_p), - nn.Linear(64, self.model_config.scalar_output_size) # 6 scalar variables - ) + # Scalar output head (6 variables, optional) + if self.include_scalar and self.model_config.scalar_output_size > 0: + self.scalar_head = nn.Sequential( + nn.Linear(self.token_dim, 64), + nn.BatchNorm1d(64), # Add BatchNorm + nn.ReLU(), + nn.Dropout(self.dropout_p), + nn.Linear(64, self.model_config.scalar_output_size) # 6 scalar variables + ) + else: + self.scalar_head = None # 2D output head (dynamic number of 2D soil variables) n_2d_vars = len(self.data_info.get('y_list_columns_2d', [])) self.matrix_head = nn.Sequential( @@ -559,7 +569,7 @@ def forward(self, time_series_data: torch.Tensor, static_data: torch.Tensor, pft_param_features = self.cnn_pft_param(x) features.append(pft_param_features) # Scalar encoder - if self.fc_scalar is not None: + if self.include_scalar and self.fc_scalar is not None: scalar_features = self.fc_scalar(scalar) # print("NaNs in scalar_features:", torch.isnan(scalar_features).sum().item(), "shape:", scalar_features.shape) features.append(scalar_features) @@ -593,9 +603,10 @@ def forward(self, time_series_data: torch.Tensor, static_data: torch.Tensor, # print("Max/Min/Mean fused_features:", fused_features.max().item(), fused_features.min().item(), fused_features.mean().item()) # Output heads outputs = {} - scalar_pred = self.scalar_head(fused_features) - # Apply non-negativity constraint to all outputs (all are pools) - outputs['scalar'] = torch.relu(scalar_pred) + if self.include_scalar and self.scalar_head is not None: + scalar_pred = self.scalar_head(fused_features) + # Apply non-negativity constraint to all outputs (all are pools) + outputs['scalar'] = torch.relu(scalar_pred) # Process PFT 1D outputs to apply specific constraints per variable pft_1d_raw_output = self.pft_1d_head(fused_features) @@ -634,7 +645,7 @@ def get_loss_weights(self) -> Dict[str, float]: """Get loss weights for different output types.""" if self.use_learnable_loss_weights: weights = {} - if self.log_sigma_scalar is not None: + if self.include_scalar and self.log_sigma_scalar is not None: weights['scalar'] = (1 / (2 * torch.exp(self.log_sigma_scalar) ** 2)).item() if self.log_sigma_matrix is not None: weights['matrix'] = (1 / (2 * torch.exp(self.log_sigma_matrix) ** 2)).item() @@ -645,9 +656,10 @@ def get_loss_weights(self) -> Dict[str, float]: return weights else: weights = { - 'scalar': 1.0, 'matrix': 1.0 } + if self.include_scalar: + weights['scalar'] = 1.0 if self.include_water: weights['water'] = 1.0 # Optionally add pft_1d if used in loss diff --git a/train_cnp_model.py b/train_cnp_model.py index 156742d..41c2751 100644 --- a/train_cnp_model.py +++ b/train_cnp_model.py @@ -269,6 +269,12 @@ def main(): help='Disable masking of absent PFTs' ) parser.set_defaults(mask_absent_pfts=True) + parser.add_argument( + '--drop-scalar-variables', + dest='drop_scalar_variables', + action='store_true', + help='Drop scalar variables from both input and output (default: include scalar variables)' + ) args = parser.parse_args() @@ -300,6 +306,10 @@ def main(): # include_water = args.with_water logger.info(f"Water variables included: {include_water}") + # Scalar variables inclusion + include_scalar = not args.drop_scalar_variables + logger.info(f"Scalar variables included: {include_scalar}") + # Get configuration (support variable list and optional model-config overrides) from config.training_config import get_cnp_combined_config config = get_cnp_combined_config( @@ -308,6 +318,7 @@ def main(): use_tva4km=args.use_tva4km, max_files=args.max_files, include_water=include_water, + include_scalar=include_scalar, variable_list_path=args.variable_list, model_config_path=args.model_config ) @@ -609,6 +620,7 @@ def main(): config.model_config, data_info, include_water=include_water, + include_scalar=include_scalar, use_learnable_loss_weights=config.training_config.use_learnable_loss_weights ) @@ -694,6 +706,7 @@ def main(): config_dict = { 'include_water': include_water, + 'include_scalar': include_scalar, 'normalization_method': args.normalization, 'data_info': data_info, 'data_counts': group_counts, diff --git a/training/trainer.py b/training/trainer.py index 845714e..0714735 100644 --- a/training/trainer.py +++ b/training/trainer.py @@ -295,15 +295,13 @@ def train_epoch(self) -> float: self.train_data['time_series'], self.train_data['static'], self.train_data['pft_param'], - self.train_data['scalar'], self.train_data['variables_1d_pft'], self.train_data['variables_2d_soil'], - self.train_data['y_scalar'], self.train_data['y_pft_1d'], self.train_data['y_soil_2d'] ] - tensor_names = ['time_series', 'static', 'pft_param', 'scalar', 'variables_1d_pft', 'variables_2d_soil', 'y_scalar', 'y_pft_1d', 'y_soil_2d'] + tensor_names = ['time_series', 'static', 'pft_param', 'variables_1d_pft', 'variables_2d_soil', 'y_pft_1d', 'y_soil_2d'] batch_sizes = [t.shape[0] for t in tensors_to_check] logger.info(f"Training tensor batch sizes: {dict(zip(tensor_names, batch_sizes))}") @@ -312,6 +310,13 @@ def train_epoch(self) -> float: logger.error(f"Tensor batch sizes are inconsistent: {dict(zip(tensor_names, batch_sizes))}") raise ValueError(f"Tensor batch sizes must be consistent. Found: {dict(zip(tensor_names, batch_sizes))}") + # Add scalar to tensors_to_check and tensor_names if present + if 'scalar' in self.train_data: + tensors_to_check.append(self.train_data['scalar']) + tensor_names.append('scalar') + if 'y_scalar' in self.train_data: + tensors_to_check.append(self.train_data['y_scalar']) + tensor_names.append('y_scalar') # Add water to tensors_to_check and tensor_names if present if 'water' in self.train_data: tensors_to_check.append(self.train_data['water']) @@ -321,35 +326,32 @@ def train_epoch(self) -> float: tensor_names.append('y_water') # Create data loader with GPU optimizations + # Build dataset based on presence of scalar and water + dataset_tensors = [ + self.train_data['time_series'], + self.train_data['static'], + self.train_data['pft_param'], + ] + if 'scalar' in self.train_data: + dataset_tensors.append(self.train_data['scalar']) + dataset_tensors.extend([ + self.train_data['variables_1d_pft'], + self.train_data['variables_2d_soil'], + ]) + if 'y_scalar' in self.train_data: + dataset_tensors.append(self.train_data['y_scalar']) + dataset_tensors.extend([ + self.train_data['y_pft_1d'], + self.train_data['y_soil_2d'], + ]) if 'water' in self.train_data and 'y_water' in self.train_data: - train_dataset = TensorDataset( - self.train_data['time_series'], - self.train_data['static'], - self.train_data['pft_param'], - self.train_data['scalar'], - self.train_data['variables_1d_pft'], - self.train_data['variables_2d_soil'], - self.train_data['y_pft_1d'], - self.train_data['y_scalar'], - self.train_data['y_soil_2d'], + dataset_tensors.extend([ self.train_data['water'], self.train_data['y_water'], - *( (self.train_data['pft_presence_mask'],) if 'pft_presence_mask' in self.train_data else () ) - ) - else: - train_dataset = TensorDataset( - self.train_data['time_series'], - self.train_data['static'], - self.train_data['pft_param'], - self.train_data['scalar'], - self.train_data['variables_1d_pft'], - self.train_data['variables_2d_soil'], - self.train_data['y_scalar'], - self.train_data['y_pft_1d'], - self.train_data['y_soil_2d'], - # Optional mask as final feature; if absent, a placeholder will be injected in-loop - *( (self.train_data['pft_presence_mask'],) if 'pft_presence_mask' in self.train_data else () ) - ) + ]) + if 'pft_presence_mask' in self.train_data: + dataset_tensors.append(self.train_data['pft_presence_mask']) + train_dataset = TensorDataset(*dataset_tensors) train_loader = DataLoader( train_dataset, @@ -367,16 +369,20 @@ def get_loss_value(loss): return loss.item() if hasattr(loss, 'item') else loss for batch_idx, batch in enumerate(progress_bar): - if 'water' in self.train_data and 'y_water' in self.train_data: - if 'pft_presence_mask' in self.train_data: - (time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, y_scalar, y_pft_1d, y_soil_2d, water, y_water, pft_presence_mask) = batch - else: - (time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, y_scalar, y_pft_1d, y_soil_2d, water, y_water) = batch - else: - if 'pft_presence_mask' in self.train_data: - (time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, y_scalar, y_pft_1d, y_soil_2d, pft_presence_mask) = batch - else: - (time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, y_scalar, y_pft_1d, y_soil_2d) = batch + # Unpack batch based on presence of scalar, water, and pft_presence_mask + idx = 0 + time_series = batch[idx]; idx += 1 + static = batch[idx]; idx += 1 + pft_param = batch[idx]; idx += 1 + scalar = batch[idx] if 'scalar' in self.train_data else None; idx += 1 if 'scalar' in self.train_data else 0 + variables_1d_pft = batch[idx]; idx += 1 + variables_2d_soil = batch[idx]; idx += 1 + y_scalar = batch[idx] if 'y_scalar' in self.train_data else None; idx += 1 if 'y_scalar' in self.train_data else 0 + y_pft_1d = batch[idx]; idx += 1 + y_soil_2d = batch[idx]; idx += 1 + water = batch[idx] if 'water' in self.train_data and 'y_water' in self.train_data else None; idx += 1 if 'water' in self.train_data and 'y_water' in self.train_data else 0 + y_water = batch[idx] if 'water' in self.train_data and 'y_water' in self.train_data else None; idx += 1 if 'water' in self.train_data and 'y_water' in self.train_data else 0 + pft_presence_mask = batch[idx] if 'pft_presence_mask' in self.train_data else None # --- DEBUG: Print tensor shapes and device before model call --- # print(f"[DEBUG] Batch {batch_idx} tensor shapes and device:") # print(f" time_series: {time_series.shape}, device: {time_series.device}") @@ -397,8 +403,10 @@ def get_loss_value(loss): variables_1d_pft = variables_1d_pft.to(self.device, non_blocking=True).contiguous() variables_2d_soil = variables_2d_soil.to(self.device, non_blocking=True).contiguous() pft_param = pft_param.to(self.device, non_blocking=True).contiguous() - scalar = scalar.to(self.device, non_blocking=True).contiguous() - y_scalar = y_scalar.to(self.device, non_blocking=True).contiguous() + if scalar is not None: + scalar = scalar.to(self.device, non_blocking=True).contiguous() + if y_scalar is not None: + y_scalar = y_scalar.to(self.device, non_blocking=True).contiguous() y_pft_1d = y_pft_1d.to(self.device, non_blocking=True).contiguous() y_soil_2d = y_soil_2d.to(self.device, non_blocking=True).contiguous() if 'water' in self.train_data and 'y_water' in self.train_data: @@ -418,12 +426,12 @@ def get_loss_value(loss): # Forward pass (with or without mixed precision) if self.use_amp and self.scaler is not None: with torch.cuda.amp.autocast(): - if 'water' in self.train_data and 'y_water' in self.train_data: + if water is not None: outputs = self.model(time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, water) else: outputs = self.model(time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil) else: - if 'water' in self.train_data and 'y_water' in self.train_data: + if water is not None: outputs = self.model(time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, water) else: outputs = self.model(time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil) @@ -445,26 +453,28 @@ def get_loss_value(loss): pass # Compute loss with variable-specific weights for scalar variables - if self.use_variable_weights and hasattr(self, 'scalar_var_weights') and self.scalar_var_weights: - # Apply variable-specific weights to scalar variables - scalar_loss = 0.0 - scalar_pred = outputs['scalar'] - - # Get variable names - scalar_vars = self.data_info.get('x_list_scalar_columns', []) - - for i, var_name in enumerate(scalar_vars): - if i < scalar_pred.size(1): # Ensure index is within bounds - var_weight = self.scalar_var_weights.get(var_name, 1.0) - var_loss = self._compute_loss(scalar_pred[:, i:i+1], y_scalar[:, i:i+1]) - scalar_loss += var_weight * var_loss - - # Normalize by number of variables to maintain scale - scalar_loss = scalar_loss / max(1, len(scalar_vars)) - loss = self.scalar_loss_weight * scalar_loss - else: - # Use standard loss calculation - loss = self.scalar_loss_weight * self._compute_loss(outputs['scalar'], y_scalar) + loss = 0.0 + if 'scalar' in outputs and y_scalar is not None: + if self.use_variable_weights and hasattr(self, 'scalar_var_weights') and self.scalar_var_weights: + # Apply variable-specific weights to scalar variables + scalar_loss = 0.0 + scalar_pred = outputs['scalar'] + + # Get variable names + scalar_vars = self.data_info.get('x_list_scalar_columns', []) + + for i, var_name in enumerate(scalar_vars): + if i < scalar_pred.size(1): # Ensure index is within bounds + var_weight = self.scalar_var_weights.get(var_name, 1.0) + var_loss = self._compute_loss(scalar_pred[:, i:i+1], y_scalar[:, i:i+1]) + scalar_loss += var_weight * var_loss + + # Normalize by number of variables to maintain scale + scalar_loss = scalar_loss / max(1, len(scalar_vars)) + loss = self.scalar_loss_weight * scalar_loss + else: + # Use standard loss calculation + loss = self.scalar_loss_weight * self._compute_loss(outputs['scalar'], y_scalar) # Vector (PFT1D): Apply differential weighting specifically to xsmrpool vector_pred = outputs['pft_1d'] # shape: [batch, n_vars*n_pfts] or [batch, n_vars, n_pfts] @@ -668,12 +678,10 @@ def validate_epoch(self) -> float: self.test_data['variables_2d_soil'], self.test_data['y_soil_2d'], self.test_data['pft_param'], - self.test_data['scalar'], - self.test_data['y_pft_1d'], - self.test_data['y_scalar'] + self.test_data['y_pft_1d'] ] - tensor_names = ['time_series', 'static', 'variables_1d_pft', 'variables_2d_soil', 'y_soil_2d', 'pft_param', 'scalar', 'y_pft_1d', 'y_scalar'] + tensor_names = ['time_series', 'static', 'variables_1d_pft', 'variables_2d_soil', 'y_soil_2d', 'pft_param', 'y_pft_1d'] batch_sizes = [t.shape[0] for t in tensors_to_check] logger.info(f"Validation tensor batch sizes: {dict(zip(tensor_names, batch_sizes))}") @@ -687,6 +695,13 @@ def validate_epoch(self) -> float: logger.error(f"Validation tensor batch sizes are inconsistent: {dict(zip(tensor_names, batch_sizes))}") raise ValueError(f"Validation tensor batch sizes must be consistent. Found: {dict(zip(tensor_names, batch_sizes))}") + # Add scalar to tensors_to_check and tensor_names if present + if 'scalar' in self.test_data: + tensors_to_check.append(self.test_data['scalar']) + tensor_names.append('scalar') + if 'y_scalar' in self.test_data: + tensors_to_check.append(self.test_data['y_scalar']) + tensor_names.append('y_scalar') # Add water to tensors_to_check and tensor_names if present if 'water' in self.test_data: tensors_to_check.append(self.test_data['water']) @@ -696,34 +711,32 @@ def validate_epoch(self) -> float: tensor_names.append('y_water') # Create data loader with GPU optimizations + # Build dataset based on presence of scalar and water + dataset_tensors = [ + self.test_data['time_series'], + self.test_data['static'], + self.test_data['pft_param'], + ] + if 'scalar' in self.test_data: + dataset_tensors.append(self.test_data['scalar']) + dataset_tensors.extend([ + self.test_data['variables_1d_pft'], + self.test_data['variables_2d_soil'], + ]) + if 'y_scalar' in self.test_data: + dataset_tensors.append(self.test_data['y_scalar']) + dataset_tensors.extend([ + self.test_data['y_pft_1d'], + self.test_data['y_soil_2d'], + ]) if 'water' in self.test_data and 'y_water' in self.test_data: - val_dataset = TensorDataset( - self.test_data['time_series'], - self.test_data['static'], - self.test_data['pft_param'], - self.test_data['scalar'], - self.test_data['variables_1d_pft'], - self.test_data['variables_2d_soil'], - self.test_data['y_scalar'], - self.test_data['y_pft_1d'], - self.test_data['y_soil_2d'], + dataset_tensors.extend([ self.test_data['water'], self.test_data['y_water'], - *( (self.test_data['pft_presence_mask'],) if 'pft_presence_mask' in self.test_data else () ) - ) - else: - val_dataset = TensorDataset( - self.test_data['time_series'], - self.test_data['static'], - self.test_data['pft_param'], - self.test_data['scalar'], - self.test_data['variables_1d_pft'], - self.test_data['variables_2d_soil'], - self.test_data['y_scalar'], - self.test_data['y_pft_1d'], - self.test_data['y_soil_2d'], - *( (self.test_data['pft_presence_mask'],) if 'pft_presence_mask' in self.test_data else () ) - ) + ]) + if 'pft_presence_mask' in self.test_data: + dataset_tensors.append(self.test_data['pft_presence_mask']) + val_dataset = TensorDataset(*dataset_tensors) val_loader = DataLoader( val_dataset, @@ -741,30 +754,37 @@ def validate_epoch(self) -> float: def get_loss_value(loss): return loss.item() if hasattr(loss, 'item') else loss for batch_idx, batch in enumerate(progress_bar): - if 'water' in self.test_data and 'y_water' in self.test_data: - if 'pft_presence_mask' in self.test_data: - (time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, y_scalar, y_pft_1d, y_soil_2d, water, y_water, pft_presence_mask) = batch - else: - (time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, y_scalar, y_pft_1d, y_soil_2d, water, y_water) = batch - else: - if 'pft_presence_mask' in self.test_data: - (time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, y_scalar, y_pft_1d, y_soil_2d, pft_presence_mask) = batch - else: - (time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, y_scalar, y_pft_1d, y_soil_2d) = batch + # Unpack batch based on presence of scalar, water, and pft_presence_mask + idx = 0 + time_series = batch[idx]; idx += 1 + static = batch[idx]; idx += 1 + pft_param = batch[idx]; idx += 1 + scalar = batch[idx] if 'scalar' in self.test_data else None; idx += 1 if 'scalar' in self.test_data else 0 + variables_1d_pft = batch[idx]; idx += 1 + variables_2d_soil = batch[idx]; idx += 1 + y_scalar = batch[idx] if 'y_scalar' in self.test_data else None; idx += 1 if 'y_scalar' in self.test_data else 0 + y_pft_1d = batch[idx]; idx += 1 + y_soil_2d = batch[idx]; idx += 1 + water = batch[idx] if 'water' in self.test_data and 'y_water' in self.test_data else None; idx += 1 if 'water' in self.test_data and 'y_water' in self.test_data else 0 + y_water = batch[idx] if 'water' in self.test_data and 'y_water' in self.test_data else None; idx += 1 if 'water' in self.test_data and 'y_water' in self.test_data else 0 + pft_presence_mask = batch[idx] if 'pft_presence_mask' in self.test_data else None # Move data to device and ensure contiguous time_series = time_series.to(self.device, non_blocking=True).contiguous() static = static.to(self.device, non_blocking=True).contiguous() variables_1d_pft = variables_1d_pft.to(self.device, non_blocking=True).contiguous() variables_2d_soil = variables_2d_soil.to(self.device, non_blocking=True).contiguous() pft_param = pft_param.to(self.device, non_blocking=True).contiguous() - scalar = scalar.to(self.device, non_blocking=True).contiguous() - y_scalar = y_scalar.to(self.device, non_blocking=True).contiguous() + if scalar is not None: + scalar = scalar.to(self.device, non_blocking=True).contiguous() + if y_scalar is not None: + y_scalar = y_scalar.to(self.device, non_blocking=True).contiguous() y_pft_1d = y_pft_1d.to(self.device, non_blocking=True).contiguous() y_soil_2d = y_soil_2d.to(self.device, non_blocking=True).contiguous() - if 'water' in self.test_data and 'y_water' in self.test_data: + if water is not None: water = water.to(self.device, non_blocking=True).contiguous() + if y_water is not None: y_water = y_water.to(self.device, non_blocking=True).contiguous() - if 'pft_presence_mask' in self.test_data: + if pft_presence_mask is not None: pft_presence_mask = pft_presence_mask.to(self.device, non_blocking=True).contiguous() # print(f"[DEBUG] variables_1d_pft shape before model (val): {variables_1d_pft.shape}") @@ -802,7 +822,9 @@ def get_loss_value(loss): pass # Compute loss - loss = self._compute_loss(outputs['scalar'], y_scalar) + loss = 0.0 + if 'scalar' in outputs and y_scalar is not None: + loss = self._compute_loss(outputs['scalar'], y_scalar) # Vector (PFT1D): base MSE vector_pred = outputs['pft_1d'] vector_targ = y_pft_1d @@ -848,10 +870,12 @@ def _prepare_batch_data(self, data: Dict[str, Any], start_idx: int, end_idx: int batch_data['time_series'] = data['time_series'][start_idx:end_idx] batch_data['static'] = data['static'][start_idx:end_idx] batch_data['pft_param'] = data['pft_param'][start_idx:end_idx] - batch_data['scalar'] = data['scalar'][start_idx:end_idx] + if 'scalar' in data: + batch_data['scalar'] = data['scalar'][start_idx:end_idx] batch_data['variables_1d_pft'] = data['variables_1d_pft'][start_idx:end_idx] batch_data['variables_2d_soil'] = data['variables_2d_soil'][start_idx:end_idx] - batch_data['y_scalar'] = data['y_scalar'][start_idx:end_idx] + if 'y_scalar' in data: + batch_data['y_scalar'] = data['y_scalar'][start_idx:end_idx] batch_data['y_pft_1d'] = data['y_pft_1d'][start_idx:end_idx] batch_data['y_soil_2d'] = data['y_soil_2d'][start_idx:end_idx] @@ -1078,10 +1102,8 @@ def evaluate(self) -> Tuple[Dict[str, Any], Dict[str, float]]: self.test_data['time_series'].shape[0], self.test_data['static'].shape[0], self.test_data['pft_param'].shape[0], - self.test_data['scalar'].shape[0], self.test_data['variables_1d_pft'].shape[0], self.test_data['variables_2d_soil'].shape[0], - self.test_data['y_scalar'].shape[0], self.test_data['y_pft_1d'].shape[0], self.test_data['y_soil_2d'].shape[0] ] @@ -1103,31 +1125,27 @@ def evaluate(self) -> Tuple[Dict[str, Any], Dict[str, float]]: return empty_predictions, default_metrics # Create evaluation data loader (optionally include presence mask) + # Build dataset based on presence of scalar + dataset_tensors = [ + self.test_data['time_series'], + self.test_data['static'], + self.test_data['pft_param'], + ] + if 'scalar' in self.test_data: + dataset_tensors.append(self.test_data['scalar']) + dataset_tensors.extend([ + self.test_data['variables_1d_pft'], + self.test_data['variables_2d_soil'], + ]) + if 'y_scalar' in self.test_data: + dataset_tensors.append(self.test_data['y_scalar']) + dataset_tensors.extend([ + self.test_data['y_pft_1d'], + self.test_data['y_soil_2d'], + ]) if 'pft_presence_mask' in self.test_data: - eval_dataset = TensorDataset( - self.test_data['time_series'], - self.test_data['static'], - self.test_data['pft_param'], - self.test_data['scalar'], - self.test_data['variables_1d_pft'], - self.test_data['variables_2d_soil'], - self.test_data['y_scalar'], - self.test_data['y_pft_1d'], - self.test_data['y_soil_2d'], - self.test_data['pft_presence_mask'] - ) - else: - eval_dataset = TensorDataset( - self.test_data['time_series'], - self.test_data['static'], - self.test_data['pft_param'], - self.test_data['scalar'], - self.test_data['variables_1d_pft'], - self.test_data['variables_2d_soil'], - self.test_data['y_scalar'], - self.test_data['y_pft_1d'], - self.test_data['y_soil_2d'] - ) + dataset_tensors.append(self.test_data['pft_presence_mask']) + eval_dataset = TensorDataset(*dataset_tensors) eval_loader = DataLoader( eval_dataset, batch_size=self.config.batch_size, @@ -1137,41 +1155,58 @@ def evaluate(self) -> Tuple[Dict[str, Any], Dict[str, float]]: num_workers=self.config.num_workers ) all_predictions = { - 'scalar': [], 'pft_1d': [], 'soil_2d': [] } all_targets = { - 'y_scalar': [], 'y_pft_1d': [], 'y_soil_2d': [] } + if 'scalar' in self.test_data: + all_predictions['scalar'] = [] + if 'y_scalar' in self.test_data: + all_targets['y_scalar'] = [] with torch.no_grad(): for batch in eval_loader: - if 'pft_presence_mask' in self.test_data: - (time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, y_scalar, y_pft_1d, y_soil_2d, pft_presence_mask) = batch - else: - (time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil, y_scalar, y_pft_1d, y_soil_2d) = batch + idx = 0 + time_series = batch[idx]; idx += 1 + static = batch[idx]; idx += 1 + pft_param = batch[idx]; idx += 1 + scalar = batch[idx] if 'scalar' in self.test_data else None; idx += 1 if 'scalar' in self.test_data else 0 + variables_1d_pft = batch[idx]; idx += 1 + variables_2d_soil = batch[idx]; idx += 1 + y_scalar = batch[idx] if 'y_scalar' in self.test_data else None; idx += 1 if 'y_scalar' in self.test_data else 0 + y_pft_1d = batch[idx]; idx += 1 + y_soil_2d = batch[idx]; idx += 1 + pft_presence_mask = batch[idx] if 'pft_presence_mask' in self.test_data else None; idx += 1 if 'pft_presence_mask' in self.test_data else 0 # Move to device time_series = time_series.to(self.device, non_blocking=True) static = static.to(self.device, non_blocking=True) pft_param = pft_param.to(self.device, non_blocking=True) - scalar = scalar.to(self.device, non_blocking=True) + if scalar is not None: + scalar = scalar.to(self.device, non_blocking=True) variables_1d_pft = variables_1d_pft.to(self.device, non_blocking=True) variables_2d_soil = variables_2d_soil.to(self.device, non_blocking=True) - y_scalar = y_scalar.to(self.device, non_blocking=True) + if y_scalar is not None: + y_scalar = y_scalar.to(self.device, non_blocking=True) y_pft_1d = y_pft_1d.to(self.device, non_blocking=True) y_soil_2d = y_soil_2d.to(self.device, non_blocking=True) - if 'pft_presence_mask' in self.test_data: + if pft_presence_mask is not None: pft_presence_mask = pft_presence_mask.to(self.device, non_blocking=True) # Forward pass if self.use_amp and self.scaler is not None: with torch.amp.autocast('cuda'): - outputs = self.model(time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil) + if scalar is not None: + outputs = self.model(time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil) + else: + outputs = self.model(time_series, static, pft_param, variables_1d_pft, variables_2d_soil) else: - outputs = self.model(time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil) + if scalar is not None: + outputs = self.model(time_series, static, pft_param, scalar, variables_1d_pft, variables_2d_soil) + else: + outputs = self.model(time_series, static, pft_param, variables_1d_pft, variables_2d_soil) # Apply presence mask to predictions if enabled - if getattr(self.config, 'mask_absent_pfts', False) and 'pft_1d' in outputs and 'pft_presence_mask' in self.test_data: + if getattr(self.config, 'mask_absent_pfts', False) and 'pft_1d' in outputs and pft_presence_mask is not None: try: vec = outputs['pft_1d'] # Determine n_vars and reshape @@ -1186,15 +1221,24 @@ def evaluate(self) -> Tuple[Dict[str, Any], Dict[str, float]]: outputs['pft_1d'] = vec.view(vec.size(0), -1) except Exception: pass - all_predictions['scalar'].append(outputs['scalar'].cpu()) + if 'scalar' in outputs: + all_predictions['scalar'].append(outputs['scalar'].cpu()) all_predictions['pft_1d'].append(outputs['pft_1d'].cpu()) all_predictions['soil_2d'].append(outputs['soil_2d'].cpu()) - all_targets['y_scalar'].append(y_scalar.cpu()) + if y_scalar is not None: + all_targets['y_scalar'].append(y_scalar.cpu()) all_targets['y_pft_1d'].append(y_pft_1d.cpu()) all_targets['y_soil_2d'].append(y_soil_2d.cpu()) - # Concatenate all batches - predictions = {k: torch.cat(v, dim=0) for k, v in all_predictions.items()} - targets = {k: torch.cat(v, dim=0) for k, v in all_targets.items()} + # Concatenate all batches (only if list is not empty) + predictions = {} + targets = {} + for k, v in all_predictions.items(): + if v: + predictions[k] = torch.cat(v, dim=0) + # Don't add empty tensors to predictions/targets - they will be handled in _calculate_metrics + for k, v in all_targets.items(): + if v: + targets[k] = torch.cat(v, dim=0) # Debug: Print shapes to identify the issue print(f"[DEBUG] Evaluation - predictions shapes:") @@ -1211,99 +1255,129 @@ def evaluate(self) -> Tuple[Dict[str, Any], Dict[str, float]]: def _calculate_metrics(self, predictions: Dict[str, torch.Tensor], targets: Dict[str, torch.Tensor]) -> Dict[str, float]: """Calculate evaluation metrics.""" metrics = {} - # Scalar - pred_scalar_full = predictions['scalar'].cpu().numpy() - target_scalar_full = targets['y_scalar'].cpu().numpy() - # Overall scalar metrics - mask_all = ~np.isnan(pred_scalar_full) & ~np.isnan(target_scalar_full) - # Ensure arrays have the same shape before flattening - if pred_scalar_full.shape != target_scalar_full.shape: - print(f"Warning: Shape mismatch - pred: {pred_scalar_full.shape}, target: {target_scalar_full.shape}") - # Use the minimum shape to avoid indexing errors - min_shape = (min(pred_scalar_full.shape[0], target_scalar_full.shape[0]), - min(pred_scalar_full.shape[1], target_scalar_full.shape[1])) - pred_scalar_full = pred_scalar_full[:min_shape[0], :min_shape[1]] - target_scalar_full = target_scalar_full[:min_shape[0], :min_shape[1]] - mask_all = ~np.isnan(pred_scalar_full) & ~np.isnan(target_scalar_full) - - # Flatten arrays and mask for overall metrics - pred_flat = pred_scalar_full.flatten() - target_flat = target_scalar_full.flatten() - mask_flat = mask_all.flatten() - mse_scalar_all = mean_squared_error(target_flat[mask_flat], pred_flat[mask_flat]) - metrics['scalar_rmse'] = np.sqrt(mse_scalar_all) - metrics['scalar_mse'] = mse_scalar_all - # Per-scalar metrics with names - num_scalar = pred_scalar_full.shape[1] - scalar_names = self.data_info.get('y_list_scalar_columns', [f'scalar_{i}' for i in range(num_scalar)]) - if len(scalar_names) != num_scalar: - scalar_names = [f'scalar_{i}' for i in range(num_scalar)] - for s in range(num_scalar): - pred_s = pred_scalar_full[:, s] - targ_s = target_scalar_full[:, s] - mask_s = ~np.isnan(pred_s) & ~np.isnan(targ_s) - if np.sum(mask_s) > 0: - mse_s = mean_squared_error(targ_s[mask_s], pred_s[mask_s]) - rmse_s = np.sqrt(mse_s) - mean_s = np.mean(targ_s[mask_s]) - nrmse_s = rmse_s / mean_s if mean_s != 0 else float('inf') - r2_s = r2_score(targ_s[mask_s], pred_s[mask_s]) - name_s = scalar_names[s] - metrics[f'{name_s}_mse'] = mse_s - metrics[f'{name_s}_rmse'] = rmse_s - metrics[f'{name_s}_nrmse'] = nrmse_s - metrics[f'{name_s}_r2'] = r2_s + # Scalar (only if present and has valid shape) + if 'scalar' in predictions and 'y_scalar' in targets: + pred_scalar = predictions['scalar'] + target_scalar = targets['y_scalar'] + # Check if tensors are not empty and have correct dimensions + if pred_scalar.numel() > 0 and target_scalar.numel() > 0 and pred_scalar.ndim >= 2 and target_scalar.ndim >= 2: + pred_scalar_full = pred_scalar.cpu().numpy() + target_scalar_full = target_scalar.cpu().numpy() + # Overall scalar metrics + mask_all = ~np.isnan(pred_scalar_full) & ~np.isnan(target_scalar_full) + # Ensure arrays have the same shape before flattening + if pred_scalar_full.shape != target_scalar_full.shape: + print(f"Warning: Shape mismatch - pred: {pred_scalar_full.shape}, target: {target_scalar_full.shape}") + # Use the minimum shape to avoid indexing errors + if len(pred_scalar_full.shape) >= 2 and len(target_scalar_full.shape) >= 2: + min_shape = (min(pred_scalar_full.shape[0], target_scalar_full.shape[0]), + min(pred_scalar_full.shape[1], target_scalar_full.shape[1])) + pred_scalar_full = pred_scalar_full[:min_shape[0], :min_shape[1]] + target_scalar_full = target_scalar_full[:min_shape[0], :min_shape[1]] + mask_all = ~np.isnan(pred_scalar_full) & ~np.isnan(target_scalar_full) + + # Flatten arrays and mask for overall metrics + pred_flat = pred_scalar_full.flatten() + target_flat = target_scalar_full.flatten() + mask_flat = mask_all.flatten() + if np.sum(mask_flat) > 0: + mse_scalar_all = mean_squared_error(target_flat[mask_flat], pred_flat[mask_flat]) + metrics['scalar_rmse'] = np.sqrt(mse_scalar_all) + metrics['scalar_mse'] = mse_scalar_all + # Per-scalar metrics with names + if len(pred_scalar_full.shape) >= 2: + num_scalar = pred_scalar_full.shape[1] + scalar_names = self.data_info.get('y_list_scalar_columns', [f'scalar_{i}' for i in range(num_scalar)]) + if len(scalar_names) != num_scalar: + scalar_names = [f'scalar_{i}' for i in range(num_scalar)] + for s in range(num_scalar): + pred_s = pred_scalar_full[:, s] + targ_s = target_scalar_full[:, s] + mask_s = ~np.isnan(pred_s) & ~np.isnan(targ_s) + if np.sum(mask_s) > 0: + mse_s = mean_squared_error(targ_s[mask_s], pred_s[mask_s]) + rmse_s = np.sqrt(mse_s) + mean_s = np.mean(targ_s[mask_s]) + nrmse_s = rmse_s / mean_s if mean_s != 0 else float('inf') + r2_s = r2_score(targ_s[mask_s], pred_s[mask_s]) + name_s = scalar_names[s] + metrics[f'{name_s}_mse'] = mse_s + metrics[f'{name_s}_rmse'] = rmse_s + metrics[f'{name_s}_nrmse'] = nrmse_s + metrics[f'{name_s}_r2'] = r2_s + else: + metrics['scalar_rmse'] = 0.0 + metrics['scalar_mse'] = 0.0 + else: + metrics['scalar_rmse'] = 0.0 + metrics['scalar_mse'] = 0.0 + else: + metrics['scalar_rmse'] = 0.0 + metrics['scalar_mse'] = 0.0 # PFT 1D - Detailed metrics per variable and PFT pred_pft_1d = predictions['pft_1d'].cpu().numpy() target_pft_1d = targets['y_pft_1d'].cpu().numpy() - # Ensure shapes are (samples, num_variables, num_pfts) - if pred_pft_1d.ndim == 2: - pred_pft_1d = pred_pft_1d.reshape(pred_pft_1d.shape[0], -1, 16) - target_pft_1d = target_pft_1d.reshape(target_pft_1d.shape[0], -1, 16) - num_variables = pred_pft_1d.shape[1] - num_pfts = pred_pft_1d.shape[2] - # Get variable names from data_info if available - var_names = self.data_info.get('y_list_columns_1d', [f'pft_1d_var_{i}' for i in range(num_variables)]) - if len(var_names) != num_variables: - var_names = [f'pft_1d_var_{i}' for i in range(num_variables)] - # Overall metrics for pft_1d - mask = ~np.isnan(pred_pft_1d) & ~np.isnan(target_pft_1d) - # Ensure arrays have the same shape before flattening - if pred_pft_1d.shape != target_pft_1d.shape: - print(f"Warning: PFT 1D shape mismatch - pred: {pred_pft_1d.shape}, target: {target_pft_1d.shape}") - # Use the minimum shape to avoid indexing errors - min_shape = tuple(min(pred_pft_1d.shape[i], target_pft_1d.shape[i]) for i in range(len(pred_pft_1d.shape))) - pred_pft_1d = pred_pft_1d[:min_shape[0], :min_shape[1], :min_shape[2]] - target_pft_1d = target_pft_1d[:min_shape[0], :min_shape[1], :min_shape[2]] - mask = ~np.isnan(pred_pft_1d) & ~np.isnan(target_pft_1d) - - # Flatten arrays and mask for overall metrics - pred_flat = pred_pft_1d.flatten() - target_flat = target_pft_1d.flatten() - mask_flat = mask.flatten() - mse_pft_1d = mean_squared_error(target_flat[mask_flat], pred_flat[mask_flat]) - metrics['pft_1d_rmse'] = np.sqrt(mse_pft_1d) - metrics['pft_1d_mse'] = mse_pft_1d - # Detailed metrics per variable and PFT - for v in range(num_variables): - var_name = var_names[v] - for p in range(num_pfts): - mask_vp = ~np.isnan(pred_pft_1d[:, v, p]) & ~np.isnan(target_pft_1d[:, v, p]) - if np.sum(mask_vp) > 0: - mse_vp = mean_squared_error(target_pft_1d[:, v, p][mask_vp], pred_pft_1d[:, v, p][mask_vp]) - rmse_vp = np.sqrt(mse_vp) - target_mean = np.mean(target_pft_1d[:, v, p][mask_vp]) - nrmse_vp = rmse_vp / target_mean if target_mean != 0 else float('inf') - r2_vp = r2_score(target_pft_1d[:, v, p][mask_vp], pred_pft_1d[:, v, p][mask_vp]) - pft_idx = p + 1 # Use 1..16 in keys - metrics[f'{var_name}_pft{pft_idx}_mse'] = mse_vp - metrics[f'{var_name}_pft{pft_idx}_rmse'] = rmse_vp - metrics[f'{var_name}_pft{pft_idx}_nrmse'] = nrmse_vp - metrics[f'{var_name}_pft{pft_idx}_r2'] = r2_vp + # Check if tensors are valid + if pred_pft_1d.ndim == 0 or pred_pft_1d.shape[0] == 0: + metrics['pft_1d_rmse'] = 0.0 + metrics['pft_1d_mse'] = 0.0 + else: + # Ensure shapes are (samples, num_variables, num_pfts) + if pred_pft_1d.ndim == 2: + pred_pft_1d = pred_pft_1d.reshape(pred_pft_1d.shape[0], -1, 16) + target_pft_1d = target_pft_1d.reshape(target_pft_1d.shape[0], -1, 16) + if pred_pft_1d.ndim < 3: + metrics['pft_1d_rmse'] = 0.0 + metrics['pft_1d_mse'] = 0.0 + else: + num_variables = pred_pft_1d.shape[1] + num_pfts = pred_pft_1d.shape[2] + # Get variable names from data_info if available + var_names = self.data_info.get('y_list_columns_1d', [f'pft_1d_var_{i}' for i in range(num_variables)]) + if len(var_names) != num_variables: + var_names = [f'pft_1d_var_{i}' for i in range(num_variables)] + # Overall metrics for pft_1d + mask = ~np.isnan(pred_pft_1d) & ~np.isnan(target_pft_1d) + # Ensure arrays have the same shape before flattening + if pred_pft_1d.shape != target_pft_1d.shape: + print(f"Warning: PFT 1D shape mismatch - pred: {pred_pft_1d.shape}, target: {target_pft_1d.shape}") + # Use the minimum shape to avoid indexing errors + min_shape = tuple(min(pred_pft_1d.shape[i], target_pft_1d.shape[i]) for i in range(len(pred_pft_1d.shape))) + pred_pft_1d = pred_pft_1d[:min_shape[0], :min_shape[1], :min_shape[2]] + target_pft_1d = target_pft_1d[:min_shape[0], :min_shape[1], :min_shape[2]] + mask = ~np.isnan(pred_pft_1d) & ~np.isnan(target_pft_1d) + + # Flatten arrays and mask for overall metrics + pred_flat = pred_pft_1d.flatten() + target_flat = target_pft_1d.flatten() + mask_flat = mask.flatten() + mse_pft_1d = mean_squared_error(target_flat[mask_flat], pred_flat[mask_flat]) + metrics['pft_1d_rmse'] = np.sqrt(mse_pft_1d) + metrics['pft_1d_mse'] = mse_pft_1d + # Detailed metrics per variable and PFT + for v in range(num_variables): + var_name = var_names[v] + for p in range(num_pfts): + mask_vp = ~np.isnan(pred_pft_1d[:, v, p]) & ~np.isnan(target_pft_1d[:, v, p]) + if np.sum(mask_vp) > 0: + mse_vp = mean_squared_error(target_pft_1d[:, v, p][mask_vp], pred_pft_1d[:, v, p][mask_vp]) + rmse_vp = np.sqrt(mse_vp) + target_mean = np.mean(target_pft_1d[:, v, p][mask_vp]) + nrmse_vp = rmse_vp / target_mean if target_mean != 0 else float('inf') + r2_vp = r2_score(target_pft_1d[:, v, p][mask_vp], pred_pft_1d[:, v, p][mask_vp]) + pft_idx = p + 1 # Use 1..16 in keys + metrics[f'{var_name}_pft{pft_idx}_mse'] = mse_vp + metrics[f'{var_name}_pft{pft_idx}_rmse'] = rmse_vp + metrics[f'{var_name}_pft{pft_idx}_nrmse'] = nrmse_vp + metrics[f'{var_name}_pft{pft_idx}_r2'] = r2_vp # Soil 2D pred_soil_2d = predictions['soil_2d'].cpu().numpy() target_soil_2d = targets['y_soil_2d'].cpu().numpy() # Ensure shapes match (samples, variables, columns, layers) + if pred_soil_2d.ndim == 0 or pred_soil_2d.shape[0] == 0: + metrics['soil_2d_rmse'] = 0.0 + metrics['soil_2d_mse'] = 0.0 + return metrics n_samples = pred_soil_2d.shape[0] if pred_soil_2d.ndim == 2: # If 2D, assume it's flattened and reshape to match target dimensions @@ -1339,6 +1413,10 @@ def _calculate_metrics(self, predictions: Dict[str, torch.Tensor], targets: Dict metrics['soil_2d_rmse'] = np.sqrt(mse_soil_2d) metrics['soil_2d_mse'] = mse_soil_2d # Per-variable per-layer metrics (aggregated across columns) + if pred_soil_2d.ndim < 2: + metrics['soil_2d_rmse'] = 0.0 + metrics['soil_2d_mse'] = 0.0 + return metrics num_variables_soil = pred_soil_2d.shape[1] num_layers_soil = pred_soil_2d.shape[3] if pred_soil_2d.ndim == 4 else 1 soil_var_names = self.data_info.get('y_list_columns_2d', [f'soil_2d_var_{i}' for i in range(num_variables_soil)]) @@ -1447,43 +1525,57 @@ def _save_predictions(self, predictions: Dict[str, np.ndarray], predictions_dir: except Exception as _e_loc: logger.warning(f"Failed to prepare location vectors: {_e_loc}") - # Save scalar predictions with inverse transformation - predictions_scalar_np = predictions['scalar'].cpu().numpy() - - # Handle shape mismatch - only use the first scalar output if model outputs more than expected - num_expected_scalars = len(self.data_info['y_list_scalar_columns']) - if predictions_scalar_np.shape[1] > num_expected_scalars: - print(f"Warning: Model outputs {predictions_scalar_np.shape[1]} scalars but only {num_expected_scalars} expected. Using first {num_expected_scalars}.") - predictions_scalar_np = predictions_scalar_np[:, :num_expected_scalars] + # Prepare scalar column names (used by both predictions and ground truth) + num_expected_scalars = len(self.data_info.get('y_list_scalar_columns', [])) - scalar_cols = self.data_info['y_list_scalar_columns'] - - # Apply inverse transformation to convert from normalized to original units - try: - if 'y_scalar' in self.scalers and self.scalers['y_scalar'] is not None: - predictions_scalar_original = self.scalers['y_scalar'].inverse_transform(predictions_scalar_np) - logger.info("Applied inverse transformation to scalar predictions") - else: + # Save scalar predictions with inverse transformation (only if available) + if 'scalar' in predictions and predictions['scalar'].numel() > 0: + predictions_scalar_np = predictions['scalar'].cpu().numpy() + + # Handle shape mismatch - only use the first scalar output if model outputs more than expected + if num_expected_scalars > 0 and predictions_scalar_np.shape[1] > num_expected_scalars: + print(f"Warning: Model outputs {predictions_scalar_np.shape[1]} scalars but only {num_expected_scalars} expected. Using first {num_expected_scalars}.") + predictions_scalar_np = predictions_scalar_np[:, :num_expected_scalars] + + scalar_cols = self.data_info.get('y_list_scalar_columns', [f'scalar_{i}' for i in range(predictions_scalar_np.shape[1])]) + if len(scalar_cols) != predictions_scalar_np.shape[1]: + scalar_cols = [f'scalar_{i}' for i in range(predictions_scalar_np.shape[1])] + + # Apply inverse transformation to convert from normalized to original units + try: + if 'y_scalar' in self.scalers and self.scalers['y_scalar'] is not None: + predictions_scalar_original = self.scalers['y_scalar'].inverse_transform(predictions_scalar_np) + logger.info("Applied inverse transformation to scalar predictions") + else: + predictions_scalar_original = predictions_scalar_np + logger.warning("No scalar scaler found, saving normalized values") + except Exception as e: + logger.warning(f"Failed to apply inverse transformation to scalar predictions: {e}") predictions_scalar_original = predictions_scalar_np - logger.warning("No scalar scaler found, saving normalized values") - except Exception as e: - logger.warning(f"Failed to apply inverse transformation to scalar predictions: {e}") - predictions_scalar_original = predictions_scalar_np - - predictions_df = pd.DataFrame(predictions_scalar_original, columns=scalar_cols) - if longitude_values is not None and latitude_values is not None and len(longitude_values) == len(predictions_df): - predictions_df.insert(0, 'Longitude', longitude_values) - predictions_df.insert(1, 'Latitude', latitude_values) - predictions_df.to_csv(os.path.join(predictions_dir, 'predictions_scalar.csv'), index=False) + + predictions_df = pd.DataFrame(predictions_scalar_original, columns=scalar_cols) + if longitude_values is not None and latitude_values is not None and len(longitude_values) == len(predictions_df): + predictions_df.insert(0, 'Longitude', longitude_values) + predictions_df.insert(1, 'Latitude', latitude_values) + predictions_df.to_csv(os.path.join(predictions_dir, 'predictions_scalar.csv'), index=False) + logger.info("Scalar predictions saved successfully") + else: + logger.warning("Scalar predictions not available, skipping scalar prediction save") # Save ground truth scalar with inverse transformation if available if 'y_scalar' in self.test_data: ground_truth_scalar_np = self.test_data['y_scalar'].cpu().numpy() + # Get scalar column names for ground truth + scalar_cols = self.data_info.get('y_list_scalar_columns', [f'scalar_{i}' for i in range(ground_truth_scalar_np.shape[1])]) + if len(scalar_cols) != ground_truth_scalar_np.shape[1]: + scalar_cols = [f'scalar_{i}' for i in range(ground_truth_scalar_np.shape[1])] + # Handle shape mismatch - ensure ground truth matches expected scalar count - if ground_truth_scalar_np.shape[1] > num_expected_scalars: + if num_expected_scalars > 0 and ground_truth_scalar_np.shape[1] > num_expected_scalars: print(f"Warning: Ground truth has {ground_truth_scalar_np.shape[1]} scalars but only {num_expected_scalars} expected. Using first {num_expected_scalars}.") ground_truth_scalar_np = ground_truth_scalar_np[:, :num_expected_scalars] + scalar_cols = scalar_cols[:num_expected_scalars] # Apply inverse transformation to ground truth as well try: From a17ae07ce6a50bed94b8cfaf10715daa06546a19 Mon Sep 17 00:00:00 2001 From: Zhuowei Gu Date: Mon, 22 Dec 2025 21:13:36 -0500 Subject: [PATCH 2/2] update data generation script and readme file, add two more options --- scripts/training_data_generation/README.md | 2 +- .../enhanced_training_dataset.py | 913 +++++++++++++++--- 2 files changed, 765 insertions(+), 150 deletions(-) diff --git a/scripts/training_data_generation/README.md b/scripts/training_data_generation/README.md index e6b1f6e..6337a82 100644 --- a/scripts/training_data_generation/README.md +++ b/scripts/training_data_generation/README.md @@ -121,7 +121,7 @@ output/ 2. **Choose your dataset mode**: ```bash # For complete ecosystem data (291 variables) - python python_scripts/enhanced_training_dataset.py --enhanced_dataset + python python_scripts/enhanced_training_dataset.py --enhanced_dataset --use_monthly_forcing --forcing_year_range "2004-2023" # For initial conditions only (195 variables) python python_scripts/enhanced_training_dataset.py --initial_only diff --git a/scripts/training_data_generation/python_scripts/enhanced_training_dataset.py b/scripts/training_data_generation/python_scripts/enhanced_training_dataset.py index c70c84b..4003787 100644 --- a/scripts/training_data_generation/python_scripts/enhanced_training_dataset.py +++ b/scripts/training_data_generation/python_scripts/enhanced_training_dataset.py @@ -219,12 +219,67 @@ def calculate_monthly_avg(time_series, time_series_length=58400): return monthly_averages -def generate_base_dataset(variable_definitions): +def detect_grid_format(ds): + """Detect if dataset uses 1D (gridcell-based) or 2D (lat/lon grid) format""" + if 'landfrac' not in ds.variables: + return None, None, None + + landfrac_var = ds.variables['landfrac'] + landfrac_dims = landfrac_var.dimensions + + if len(landfrac_dims) == 1: + return '1d', landfrac_dims[0], None + elif len(landfrac_dims) == 2: + return '2d', landfrac_dims[0], landfrac_dims[1] + else: + return None, None, None + +def get_gridcell_value(var, var_dims, gridcell_idx, grid_format, grid_info, dim1=None, dim2=None, lat_idx=None, lon_idx=None): + """Get a scalar value from a variable, handling both 1D and 2D grid formats""" + val = None + + if grid_format == '2d' and len(var_dims) == 2 and var_dims[0] == dim1 and var_dims[1] == dim2: + # Variable is 2D (lat, lon) format + if lat_idx is not None and lon_idx is not None: + val = var[lat_idx, lon_idx] + else: + # Fallback: flatten and index + val = var[:].flatten()[gridcell_idx] + elif len(var_dims) == 1: + # Variable is 1D, use direct indexing + val = var[gridcell_idx] + else: + # Try direct indexing as fallback + try: + val = var[gridcell_idx] + except: + # Last resort: flatten and index + val = var[:].flatten()[gridcell_idx] + + # Convert MaskedArray to regular array if needed + if hasattr(val, 'data'): # MaskedArray + val = val.data + + # Ensure scalar value + if isinstance(val, np.ndarray): + if val.size == 1: + val = val.item() + elif val.size > 1: + # Take first element if array + val = val.flatten()[0] + + return val + +def generate_base_dataset(variable_definitions, use_monthly_forcing=False, forcing_year_range="1980-1999"): """Generate base training dataset (72_dataset_construction.py logic)""" print(f"\n{'='*80}") print("STEP 1: Base Dataset Generation (72_dataset_construction.py)") print(f"{'='*80}") + if use_monthly_forcing: + print("Using pre-computed monthly average forcing data") + print(f" Year range: {forcing_year_range}") + # File paths from config surface_data_files = config.surface_data_files ad_spinup_history_files = config.ad_spinup_history_files @@ -236,13 +291,18 @@ def generate_base_dataset(variable_definitions): forcing_files = {} for var_name in variable_definitions['time_series_vars']: # Look for files containing the variable name in forcing_netcdf directory - pattern = os.path.join(config.forcing_netcdf_output_dir, f'*{var_name}*1980-1999.nc') + # Support both old format (*VAR*1980-1999.nc) and new format (*VAR*2004-2023.nc) + pattern = os.path.join(config.forcing_netcdf_output_dir, f'*{var_name}*{forcing_year_range}.nc') matching_files = glob.glob(pattern) + if not matching_files: + # Try alternative pattern without year range in filename + pattern_alt = os.path.join(config.forcing_netcdf_output_dir, f'{var_name}_*.nc') + matching_files = glob.glob(pattern_alt) if matching_files: forcing_files[var_name] = matching_files[0] # Use first match - print(f"✅ Found forcing file: {os.path.basename(matching_files[0])}") + print(f"Found forcing file: {os.path.basename(matching_files[0])}") else: - print(f"⚠️ Forcing file not found for {var_name}: {pattern}") + print(f"Forcing file not found for {var_name}: {pattern}") print(f"Found {len(forcing_files)} forcing files") @@ -262,7 +322,7 @@ def generate_base_dataset(variable_definitions): ds_forcing = {} for var_name, file_path in forcing_files.items(): ds_forcing[var_name] = nc.Dataset(file_path) - print(f"✅ Forcing data loaded: {var_name}") + print(f"Forcing data loaded: {var_name}") # Load future files for Y variables ds_h0_list = [nc.Dataset(fp) for fp in final_spinup_history_files] @@ -270,28 +330,141 @@ def generate_base_dataset(variable_definitions): print(f"All files loaded in {time.time() - start_time:.2f} seconds") + # Detect grid format for history file (ds2) + # Detect grid format for both ds1 (surface) and ds2 (history) files + grid_format_ds1, dim1_ds1, dim2_ds1 = detect_grid_format(ds1) + grid_format_ds2, dim1_ds2, dim2_ds2 = detect_grid_format(ds2) + print(f"Surface file (ds1) grid format: {grid_format_ds1} (dims: {dim1_ds1}, {dim2_ds1})") + print(f"History file (ds2) grid format: {grid_format_ds2} (dims: {dim1_ds2}, {dim2_ds2})") + + # For backward compatibility, grid_format refers to ds2 format (used for history variables) + grid_format = grid_format_ds2 + dim1 = dim1_ds2 + dim2 = dim2_ds2 + # Get coordinates and build spatial filtering - lats = ds2.variables['lat'][:] - lons = ds2.variables['lon'][:] - landmask = ds2.variables['landfrac'][:] + landmask_var = ds2.variables['landfrac'] + landmask = landmask_var[:] + + # Store grid format info and conversion mappings for both ds1 and ds2 + grid_info = { + 'format': grid_format, # ds2 format (for backward compatibility) + 'format_ds1': grid_format_ds1, # ds1 format + 'format_ds2': grid_format_ds2, # ds2 format + 'dim1_ds1': dim1_ds1, + 'dim2_ds1': dim2_ds1, + 'dim1_ds2': dim1_ds2, + 'dim2_ds2': dim2_ds2, + 'lat_idx_map': {}, # flat_idx -> lat_idx for 2D format (ds2) + 'lon_idx_map': {}, # flat_idx -> lon_idx for 2D format (ds2) + 'lat_idx_map_ds1': {}, # flat_idx -> lat_idx for 2D format (ds1) + 'lon_idx_map_ds1': {}, # flat_idx -> lon_idx for 2D format (ds1) + 'n_lon': None, + 'n_lon_ds1': None + } - # Filter for land gridcells - valid_mask = (landmask > 0) - valid_gridcells = np.where(valid_mask)[0] + # Build mappings for ds2 (history file) + if grid_format_ds2 == '1d': + # 1D format: (gridcell) or (lndgrid) + lats = ds2.variables['lat'][:] + lons = ds2.variables['lon'][:] + + # Filter for land gridcells + valid_mask = (landmask > 0) + valid_gridcells = np.where(valid_mask)[0] + + # Get coordinates for valid gridcells (direct indexing) + query_coords = np.array([(lats[i], lons[i]) for i in valid_gridcells]) + + elif grid_format_ds2 == '2d': + # 2D format: (lat, lon) + lat_coords = ds2.variables['lat'][:] + lon_coords = ds2.variables['lon'][:] + grid_info['n_lon'] = len(lon_coords) + + # Filter for land gridcells (2D mask) + valid_mask = (landmask > 0) + valid_lat_indices, valid_lon_indices = np.where(valid_mask) + + # Convert 2D indices to flat gridcell indices + valid_gridcells_2d = valid_lat_indices * grid_info['n_lon'] + valid_lon_indices + + # Get coordinates for valid gridcells + query_coords_2d = np.array([(lat_coords[i], lon_coords[j]) + for i, j in zip(valid_lat_indices, valid_lon_indices)]) + + # Create mapping from flat index to 2D indices for ds2 + for i, flat_idx in enumerate(valid_gridcells_2d): + grid_info['lat_idx_map'][i] = valid_lat_indices[i] + grid_info['lon_idx_map'][i] = valid_lon_indices[i] + + # For compatibility with later code + lats = np.array([lat_coords[i] for i in valid_lat_indices]) + lons = np.array([lon_coords[j] for j in valid_lon_indices]) + query_coords = query_coords_2d + valid_gridcells = valid_gridcells_2d + else: + raise ValueError(f"Unsupported grid format: {grid_format}") + + # Build mappings for ds1 (surface file) if it's 2D format + # This needs to be done after we have gridcell coordinates from ds10 + if grid_format_ds1 == '2d' and 'landfrac' in ds1.variables: + print("Building ds1 (surface file) 2D index mappings...") + landmask_ds1 = ds1.variables['landfrac'][:] + lat_coords_ds1 = ds1.variables[dim1_ds1][:] + lon_coords_ds1 = ds1.variables[dim2_ds1][:] + grid_info['n_lon_ds1'] = len(lon_coords_ds1) + + # Filter for land gridcells in ds1 (2D mask) + valid_mask_ds1 = (landmask_ds1 > 0) + valid_lat_indices_ds1, valid_lon_indices_ds1 = np.where(valid_mask_ds1) + + # Get restart file coordinates (these are defined later, so we'll build mapping in the main loop if needed) + # For now, we'll build it after restart coordinates are available + print(" ds1 2D mapping will be built dynamically during processing") - print(f"Total land gridcells: {len(valid_gridcells)}") + print(f"Total land gridcells (before deduplication): {len(valid_gridcells)}") print(f"Latitude range: [{lats.min():.2f}, {lats.max():.2f}]") print(f"Longitude range: [{lons.min():.2f}, {lons.max():.2f}]") # Build KDTree indices print("Building KDTree indices...") - # Restart file coordinates + # Restart file coordinates (these are the unique gridcells we want - 20975 total) gridcell_lat = ds10.variables['grid1d_lat'][:] gridcell_lon = ds10.variables['grid1d_lon'][:] restart_grid_coords = np.vstack((gridcell_lat, gridcell_lon)).T restart_tree = cKDTree(restart_grid_coords) + # Build mappings for ds1 (surface file) if it's 2D format + # Now that we have restart coordinates, we can build the mapping + if grid_format_ds1 == '2d' and 'landfrac' in ds1.variables and len(grid_info['lat_idx_map_ds1']) == 0: + print("Building ds1 (surface file) 2D index mappings from restart coordinates...") + landmask_ds1 = ds1.variables['landfrac'][:] + lat_coords_ds1 = ds1.variables[dim1_ds1][:] + lon_coords_ds1 = ds1.variables[dim2_ds1][:] + grid_info['n_lon_ds1'] = len(lon_coords_ds1) + + # Filter for land gridcells in ds1 (2D mask) + valid_mask_ds1 = (landmask_ds1 > 0) + valid_lat_indices_ds1, valid_lon_indices_ds1 = np.where(valid_mask_ds1) + + # Build KDTree for ds1 coordinates + ds1_coords = np.vstack((lat_coords_ds1[valid_lat_indices_ds1], + lon_coords_ds1[valid_lon_indices_ds1])).T + ds1_tree = cKDTree(ds1_coords) + + # Query restart coordinates to find closest ds1 gridcell + _, ds1_matched_indices = ds1_tree.query(restart_grid_coords, k=1) + + # Create mapping: restart_idx -> (lat_idx_ds1, lon_idx_ds1) + for restart_idx in range(len(gridcell_lat)): + if ds1_matched_indices[restart_idx] < len(valid_lat_indices_ds1): + matched_flat_idx = ds1_matched_indices[restart_idx] + grid_info['lat_idx_map_ds1'][restart_idx] = valid_lat_indices_ds1[matched_flat_idx] + grid_info['lon_idx_map_ds1'][restart_idx] = valid_lon_indices_ds1[matched_flat_idx] + print(f" Built {len(grid_info['lat_idx_map_ds1'])} ds1 2D mappings") + # Forcing file coordinates forcing_lats = list(ds_forcing.values())[0].variables['LATIXY'][0, :] forcing_lons = list(ds_forcing.values())[0].variables['LONGXY'][0, :] @@ -299,14 +472,43 @@ def generate_base_dataset(variable_definitions): forcing_tree = cKDTree(forcing_grid_coords) # Query coordinates for valid gridcells - query_coords = np.array([(lats[i], lons[i]) for i in valid_gridcells]) _, all_restart_indices = restart_tree.query(query_coords, k=1) _, all_forcing_indices = forcing_tree.query(query_coords, k=1) - print("✅ KDTree indices built") + # Deduplicate: ensure each unique restart gridcell is only processed once + # Multiple (lat,lon) pairs may map to the same restart gridcell + seen_restart_indices = {} + unique_indices = [] + for i, restart_idx in enumerate(all_restart_indices): + if restart_idx not in seen_restart_indices: + unique_indices.append(i) + seen_restart_indices[restart_idx] = i + + # Filter to unique gridcells + valid_gridcells = valid_gridcells[unique_indices] + all_restart_indices = all_restart_indices[unique_indices] + all_forcing_indices = all_forcing_indices[unique_indices] + lats = lats[unique_indices] + lons = lons[unique_indices] + + # Update grid_info mappings to reflect deduplication + if grid_format == '2d': + old_lat_idx_map = grid_info['lat_idx_map'].copy() + old_lon_idx_map = grid_info['lon_idx_map'].copy() + grid_info['lat_idx_map'] = {new_idx: old_lat_idx_map[old_idx] + for new_idx, old_idx in enumerate(unique_indices)} + grid_info['lon_idx_map'] = {new_idx: old_lon_idx_map[old_idx] + for new_idx, old_idx in enumerate(unique_indices)} + + print(f"Total unique gridcells (after deduplication): {len(valid_gridcells)}") + print(f"Expected: 20975 gridcells") + if len(valid_gridcells) != 20975: + print(f"⚠️ Warning: Expected 20975 gridcells but got {len(valid_gridcells)}") + + print("KDTree indices built and deduplication completed") # Pre-load forcing data into memory for optimization - print("🚀 Pre-loading forcing data into memory...") + print("Pre-loading forcing data into memory...") forcing_data = {} for var_name, ds in ds_forcing.items(): forcing_data[var_name] = ds.variables[var_name][:, 0, :] # (time, 1, grid_cells) @@ -327,7 +529,7 @@ def generate_base_dataset(variable_definitions): pft_map[grid_id] = np.where(pft_gridcell_index == grid_id)[0] column_map[grid_id] = np.where(column_gridcell_index == grid_id)[0] - print("✅ Index mappings built") + print("Index mappings built") # Process data in batches batch_size = 1000 @@ -347,7 +549,7 @@ def generate_base_dataset(variable_definitions): print(f" batch_gridcells range: {batch_gridcells[0]} to {batch_gridcells[-1]}") print(f" Unique gridcells in batch: {len(set(batch_gridcells))}") if len(set(batch_gridcells)) != len(batch_gridcells): - print(f" ⚠️ WARNING: Duplicate gridcells detected in batch!") + print(f"WARNING: Duplicate gridcells detected in batch!") batch_start_time = time.time() # Initialize data dictionary dynamically - completely from CNP_IO file @@ -387,13 +589,14 @@ def generate_base_dataset(variable_definitions): data_dict[f'Y_{var_name}'] = [] # Process each gridcell in the batch - for k, gridcell_idx in enumerate(batch_gridcells): + for k in range(len(batch_gridcells)): if k % 100 == 0: - print(f" Processing gridcell {k}/{len(batch_gridcells)} (idx={gridcell_idx})") + print(f" Processing gridcell {k}/{len(batch_gridcells)}") # Get indices - restart_idx = batch_restart_indices[k] - forcing_idx = batch_forcing_indices[k] + restart_idx = int(batch_restart_indices[k]) # Index for ds1 (surface) and ds10 (restart) files + forcing_idx = int(batch_forcing_indices[k]) # Index for forcing files + gridcell_idx_history = batch_gridcells[k] # Index for ds2 (history) file - may be 2D flattened gridcell_id = restart_idx + 1 pft_indices_for_cell = pft_map.get(gridcell_id, []) @@ -409,63 +612,213 @@ def generate_base_dataset(variable_definitions): if k == 999: print(f" After 1000th gridcell: Latitude={len(data_dict['Latitude'])}, FLDS={len(data_dict.get('FLDS', []))}") - # Process forcing data (time series variables) - store raw data like original script + # Process forcing data (time series variables) for var_name in variable_definitions['time_series_vars']: if var_name in forcing_data: time_series = forcing_data[var_name][:, forcing_idx] - # Store raw time series like original script (will be processed later) - data_dict[var_name].append(time_series) + # If using monthly forcing, data is already monthly average, store as-is + # Otherwise, store raw time series (will be processed later in post-processing) + if use_monthly_forcing: + # Data is already monthly average, convert to list of floats + data_dict[var_name].append(time_series.tolist() if isinstance(time_series, np.ndarray) else list(time_series)) + else: + # Store raw time series (will be processed later) + data_dict[var_name].append(time_series) else: data_dict[var_name].append([]) # Add empty list if variable not found # Process surface properties for var_name in variable_definitions['surface_vars']: if var_name == 'Latitude': - data_dict[var_name].append(lats[gridcell_idx]) + # Use global index (start_idx + k) since lats array is already filtered for valid gridcells + global_idx = start_idx + k + data_dict[var_name].append(lats[global_idx]) elif var_name == 'Longitude': - data_dict[var_name].append(lons[gridcell_idx]) + # Use global index (start_idx + k) since lons array is already filtered for valid gridcells + global_idx = start_idx + k + data_dict[var_name].append(lons[global_idx]) elif var_name == 'landfrac': # landfrac comes from history file (ds2), not surface file (ds1) # Convert to float64 to match reference file - landfrac_val = ds2.variables['landfrac'][gridcell_idx] + global_idx = start_idx + k + if grid_format == '1d': + landfrac_val = ds2.variables['landfrac'][gridcell_idx_history] + elif grid_format == '2d': + lat_idx = grid_info['lat_idx_map'].get(global_idx, None) + lon_idx = grid_info['lon_idx_map'].get(global_idx, None) + if lat_idx is not None and lon_idx is not None: + landfrac_val = ds2.variables['landfrac'][lat_idx, lon_idx] + else: + # Fallback: use flat index + landfrac_val = ds2.variables['landfrac'].flatten()[gridcell_idx_history] + else: + landfrac_val = ds2.variables['landfrac'][gridcell_idx_history] + if hasattr(landfrac_val, 'data'): # MaskedArray landfrac_val = landfrac_val.data + # Ensure scalar value + if isinstance(landfrac_val, np.ndarray) and landfrac_val.size > 1: + landfrac_val = landfrac_val.item() if landfrac_val.size == 1 else landfrac_val[0] data_dict[var_name].append(float(landfrac_val)) elif var_name == 'PCT_CLAY': # Store PCT_CLAY as a list (all levels) - pct_clay_data = ds1.variables['PCT_CLAY'][:, gridcell_idx] - # Convert MaskedArray to regular array and ensure float64 - if hasattr(pct_clay_data, 'data'): # MaskedArray - pct_clay_data = pct_clay_data.data - data_dict[var_name].append(pct_clay_data.astype(np.float64).tolist()) + # PCT_CLAY is from ds1, check ds1's grid format and variable dimensions + var_obj = ds1.variables['PCT_CLAY'] + var_dims = var_obj.dimensions + grid_format_ds1 = grid_info['format_ds1'] + dim1_ds1 = grid_info['dim1_ds1'] + dim2_ds1 = grid_info['dim2_ds1'] + + try: + if grid_format_ds1 == '2d' and len(var_dims) == 3 and var_dims[1] == dim1_ds1 and var_dims[2] == dim2_ds1: + # Variable is (level, lat, lon) in ds1 - use 2D indexing + lat_idx = grid_info['lat_idx_map_ds1'].get(restart_idx, None) + lon_idx = grid_info['lon_idx_map_ds1'].get(restart_idx, None) + if lat_idx is not None and lon_idx is not None: + pct_clay_data = var_obj[:, lat_idx, lon_idx] + else: + # Fallback: use 1D indexing if mapping not available + pct_clay_data = var_obj[:, restart_idx] if restart_idx < var_obj.shape[-1] else var_obj[:, 0] + else: + # Variable is (level, gridcell) or 1D format - use 1D indexing with restart_idx + if len(var_dims) >= 2 and restart_idx < var_obj.shape[-1]: + pct_clay_data = var_obj[:, restart_idx] + else: + print(f" Warning: restart_idx {restart_idx} out of bounds for PCT_CLAY (shape: {var_obj.shape})") + pct_clay_data = np.zeros(var_obj.shape[0]) + + # Convert to numpy array + pct_clay_data = np.asarray(pct_clay_data) + # Convert MaskedArray to regular array and ensure float64 + if hasattr(pct_clay_data, 'data'): # MaskedArray + pct_clay_data = pct_clay_data.data + data_dict[var_name].append(pct_clay_data.astype(np.float64).tolist()) + except (IndexError, ValueError, TypeError) as e: + print(f" Error processing PCT_CLAY: {e}") + data_dict[var_name].append([]) elif var_name == 'PCT_SAND': # Store PCT_SAND as a list (all levels) - pct_sand_data = ds1.variables['PCT_SAND'][:, gridcell_idx] - # Convert MaskedArray to regular array and ensure float64 - if hasattr(pct_sand_data, 'data'): # MaskedArray - pct_sand_data = pct_sand_data.data - data_dict[var_name].append(pct_sand_data.astype(np.float64).tolist()) + # PCT_SAND is from ds1, check ds1's grid format and variable dimensions + var_obj = ds1.variables['PCT_SAND'] + var_dims = var_obj.dimensions + grid_format_ds1 = grid_info['format_ds1'] + dim1_ds1 = grid_info['dim1_ds1'] + dim2_ds1 = grid_info['dim2_ds1'] + + try: + if grid_format_ds1 == '2d' and len(var_dims) == 3 and var_dims[1] == dim1_ds1 and var_dims[2] == dim2_ds1: + # Variable is (level, lat, lon) in ds1 - use 2D indexing + lat_idx = grid_info['lat_idx_map_ds1'].get(restart_idx, None) + lon_idx = grid_info['lon_idx_map_ds1'].get(restart_idx, None) + if lat_idx is not None and lon_idx is not None: + pct_sand_data = var_obj[:, lat_idx, lon_idx] + else: + # Fallback: use 1D indexing if mapping not available + pct_sand_data = var_obj[:, restart_idx] if restart_idx < var_obj.shape[-1] else var_obj[:, 0] + else: + # Variable is (level, gridcell) or 1D format - use 1D indexing with restart_idx + if len(var_dims) >= 2 and restart_idx < var_obj.shape[-1]: + pct_sand_data = var_obj[:, restart_idx] + else: + print(f" Warning: restart_idx {restart_idx} out of bounds for PCT_SAND (shape: {var_obj.shape})") + pct_sand_data = np.zeros(var_obj.shape[0]) + + # Convert to numpy array + pct_sand_data = np.asarray(pct_sand_data) + # Convert MaskedArray to regular array and ensure float64 + if hasattr(pct_sand_data, 'data'): # MaskedArray + pct_sand_data = pct_sand_data.data + data_dict[var_name].append(pct_sand_data.astype(np.float64).tolist()) + except (IndexError, ValueError, TypeError) as e: + print(f" Error processing PCT_SAND: {e}") + data_dict[var_name].append([]) elif var_name.startswith('PCT_NAT_PFT_') or var_name.startswith('PCT_CLAY_') or var_name.startswith('PCT_SAND_'): # Handle 2D variables with level indices if '_' in var_name: level_idx = int(var_name.split('_')[-1]) base_var = '_'.join(var_name.split('_')[:-1]) # e.g., 'PCT_CLAY' if base_var in ds1.variables: - pct_val = ds1.variables[base_var][level_idx, gridcell_idx] - # Convert MaskedArray to regular array and ensure float64 - if hasattr(pct_val, 'data'): # MaskedArray - pct_val = pct_val.data - data_dict[var_name].append(float(pct_val)) + var_obj = ds1.variables[base_var] + var_dims = var_obj.dimensions + var_shape = var_obj.shape + grid_format_ds1 = grid_info['format_ds1'] + dim1_ds1 = grid_info['dim1_ds1'] + dim2_ds1 = grid_info['dim2_ds1'] + + try: + if grid_format_ds1 == '2d' and len(var_dims) == 3 and var_dims[1] == dim1_ds1 and var_dims[2] == dim2_ds1: + # Variable is (level, lat, lon) in ds1 - use 2D indexing + lat_idx = grid_info['lat_idx_map_ds1'].get(restart_idx, None) + lon_idx = grid_info['lon_idx_map_ds1'].get(restart_idx, None) + if lat_idx is not None and lon_idx is not None and level_idx < var_shape[0]: + pct_val_raw = var_obj[level_idx, lat_idx, lon_idx] + pct_val = np.asarray(pct_val_raw).item() + else: + # Fallback: use 1D indexing if mapping not available + if level_idx < var_shape[0] and restart_idx < var_shape[1]: + pct_val_raw = var_obj[level_idx, restart_idx] + pct_val = np.asarray(pct_val_raw).item() + else: + pct_val = 0.0 + else: + # Variable is (level, gridcell) or 1D format - use 1D indexing with restart_idx + if level_idx < var_shape[0] and restart_idx < var_shape[1]: + pct_val_raw = var_obj[level_idx, restart_idx] + pct_val = np.asarray(pct_val_raw).item() + else: + print(f" Warning: Index out of bounds for {var_name} (level_idx={level_idx}, restart_idx={restart_idx}, shape={var_shape})") + pct_val = 0.0 + + # Final check: ensure we have a numeric value + if not isinstance(pct_val, (int, float, np.integer, np.floating)): + pct_val = float(pct_val) + + data_dict[var_name].append(float(pct_val)) + except (IndexError, ValueError, TypeError, AttributeError) as e: + print(f" Error processing {var_name}: {e}") + data_dict[var_name].append(0.0) else: data_dict[var_name].append(0.0) else: data_dict[var_name].append(0.0) elif var_name in ds1.variables: - # 1D variables - val = ds1.variables[var_name][gridcell_idx] - # Convert MaskedArray to regular array - if hasattr(val, 'data'): # MaskedArray - val = val.data + # ds1 variables are indexed by restart_idx (gridcell index), not gridcell_idx_history + var_obj = ds1.variables[var_name] + var_dims = var_obj.dimensions + var_shape = var_obj.shape + + try: + if len(var_dims) == 1: + # Variable is 1D (gridcell) - typical for surface scalar variables + if restart_idx < var_shape[0]: + # Read value and immediately convert to scalar using np.asarray().item() + # This handles memoryview, MaskedArray, and other netCDF types + val = np.asarray(var_obj[restart_idx]).item() + else: + print(f" Error: restart_idx {restart_idx} out of bounds for {var_name} (shape: {var_shape})") + val = 0.0 + else: + # Multi-dimensional variable - surface variables should typically be 1D + # This might be incorrectly classified, but try to handle it + print(f" Warning: {var_name} has {len(var_dims)} dimensions but is in surface_vars") + # For multi-dim, assume last dimension is gridcell, take first element of other dims + if restart_idx < var_shape[-1]: + indices = [0] * (len(var_shape) - 1) + [restart_idx] + val = np.asarray(var_obj[tuple(indices)]).item() + else: + val = 0.0 + except (IndexError, ValueError, TypeError, AttributeError) as e: + print(f" Error indexing {var_name} (dims: {var_dims}, shape: {var_shape}, restart_idx: {restart_idx}): {e}") + val = 0.0 + + # Final check: ensure we have a numeric value + if not isinstance(val, (int, float, np.integer, np.floating)): + try: + val = float(val) + except (ValueError, TypeError) as e: + print(f" Error: Could not convert {var_name} to float: {val}, type: {type(val)}, error: {e}") + val = 0.0 + # Handle integer variables if var_name in ['SOIL_COLOR', 'SOIL_ORDER']: data_dict[var_name].append(int(val)) @@ -477,10 +830,72 @@ def generate_base_dataset(variable_definitions): # Process scalar variables from history file for var_name in variable_definitions['scalar_vars']: if var_name in ds2.variables: - val = ds2.variables[var_name][0, gridcell_idx] - # Convert MaskedArray to regular array - if hasattr(val, 'data'): # MaskedArray - val = val.data + var_obj = ds2.variables[var_name] + var_dims = var_obj.dimensions + + try: + # Check if variable has time dimension + if len(var_dims) >= 2 and var_dims[0] == 'time': + # Variable has (time, ...) dimensions + if grid_format == '2d' and len(var_dims) == 3 and var_dims[1] == dim1 and var_dims[2] == dim2: + # Variable is (time, lat, lon) + global_idx = start_idx + k + lat_idx = grid_info['lat_idx_map'].get(global_idx, None) + lon_idx = grid_info['lon_idx_map'].get(global_idx, None) + if lat_idx is not None and lon_idx is not None: + val_raw = var_obj[0, lat_idx, lon_idx] + val = np.asarray(val_raw).item() + else: + val_raw = var_obj[0, :, :].flatten()[gridcell_idx_history] + val = np.asarray(val_raw).item() + else: + # Variable is (time, gridcell) or similar - use gridcell_idx_history for ds2 + if grid_format == '1d' and gridcell_idx_history < var_obj.shape[1]: + val_raw = var_obj[0, gridcell_idx_history] + val = np.asarray(val_raw).item() + else: + # For 2D format but variable is not 2D, try using gridcell_idx_history + if gridcell_idx_history < var_obj.shape[1]: + val_raw = var_obj[0, gridcell_idx_history] + else: + val_raw = var_obj[0, 0] + val = np.asarray(val_raw).item() + else: + # Variable doesn't have time dimension + if grid_format == '2d' and len(var_dims) == 2 and var_dims[0] == dim1 and var_dims[1] == dim2: + global_idx = start_idx + k + lat_idx = grid_info['lat_idx_map'].get(global_idx, None) + lon_idx = grid_info['lon_idx_map'].get(global_idx, None) + if lat_idx is not None and lon_idx is not None: + val_raw = var_obj[lat_idx, lon_idx] + val = np.asarray(val_raw).item() + else: + val_raw = var_obj[:].flatten()[gridcell_idx_history] + val = np.asarray(val_raw).item() + else: + # Variable is 1D or other format - use gridcell_idx_history for ds2 + if grid_format == '1d' and gridcell_idx_history < var_obj.shape[0]: + val_raw = var_obj[gridcell_idx_history] + val = np.asarray(val_raw).item() + else: + # Try to handle gracefully + if gridcell_idx_history < var_obj.shape[0]: + val_raw = var_obj[gridcell_idx_history] + else: + val_raw = var_obj[0] + val = np.asarray(val_raw).item() + + # Final check: ensure we have a numeric value (val should already be scalar from np.asarray().item()) + if not isinstance(val, (int, float, np.integer, np.floating)): + try: + val = float(val) + except (ValueError, TypeError) as e: + print(f" Error: Could not convert {var_name} to float: {val}, type: {type(val)}, error: {e}") + val = 0.0 + except (IndexError, ValueError, TypeError, AttributeError) as e: + print(f" Error processing {var_name} (dims: {var_dims}): {e}") + val = 0.0 + # Handle integer variables if var_name in ['SOIL_COLOR', 'SOIL_ORDER']: data_dict[var_name].append(int(val)) @@ -491,7 +906,17 @@ def generate_base_dataset(variable_definitions): y_vals = [] for ds_h0 in ds_h0_list: if var_name in ds_h0.variables: - y_val = ds_h0.variables[var_name][0, gridcell_idx] + # Y variables are from final_spinup history files (same format as ds2) + if grid_format == '2d': + global_idx = start_idx + k + lat_idx = grid_info['lat_idx_map'].get(global_idx, None) + lon_idx = grid_info['lon_idx_map'].get(global_idx, None) + if lat_idx is not None and lon_idx is not None: + y_val = ds_h0.variables[var_name][0, lat_idx, lon_idx] + else: + y_val = ds_h0.variables[var_name][0, :, :].flatten()[gridcell_idx_history] + else: + y_val = ds_h0.variables[var_name][0, gridcell_idx_history] # Convert MaskedArray to regular array if hasattr(y_val, 'data'): # MaskedArray y_val = y_val.data @@ -633,7 +1058,7 @@ def generate_base_dataset(variable_definitions): lengths = {key: len(values) for key, values in data_dict.items()} unique_lengths = set(lengths.values()) if len(unique_lengths) > 1: - print(f" ❌ Length mismatch detected!") + print(f"Length mismatch detected!") for length in unique_lengths: vars_with_length = [k for k, v in lengths.items() if v == length] print(f" Length {length}: {len(vars_with_length)} variables") @@ -648,7 +1073,7 @@ def generate_base_dataset(variable_definitions): batch_files.append(batch_save_path) # Add to list for post-processing batch_time = time.time() - batch_start_time - print(f"✅ Batch {batch_number} completed: {batch_time:.2f}s") + print(f" Batch {batch_number} completed: {batch_time:.2f}s") print(f" Path: {batch_save_path}") print(f" Shape: {df_batch.shape}") print(f" Columns: {len(df_batch.columns)}") @@ -665,14 +1090,17 @@ def generate_base_dataset(variable_definitions): for ds_r in ds_r_list: ds_r.close() - print("✅ All NetCDF files closed") - print(f"✅ Base dataset generation completed!") + print("All NetCDF files closed") + print(f"Base dataset generation completed!") print(f"Total batches: {batch_number - 1}") print(f"Output directory: {output_dir}") # Post-processing like original script print(f"\n{'='*80}") - print("POST-PROCESSING: Converting to monthly averages and expanding variables") + if use_monthly_forcing: + print("POST-PROCESSING: Expanding variables (skipping monthly average conversion - data already monthly)") + else: + print("POST-PROCESSING: Converting to monthly averages and expanding variables") print(f"{'='*80}") # Process all generated files @@ -682,12 +1110,15 @@ def generate_base_dataset(variable_definitions): # Load the file df = pd.read_pickle(file_path) - # Process time series columns (convert to monthly averages) - print(f" Processing time series columns...") - for col in variable_definitions['time_series_vars']: - if col in df.columns: - print(f" Processing {col}...") - df[col] = df[col].apply(calculate_monthly_avg) + # Process time series columns (convert to monthly averages only if not using pre-computed monthly data) + if not use_monthly_forcing: + print(f" Processing time series columns (converting to monthly averages)...") + for col in variable_definitions['time_series_vars']: + if col in df.columns: + print(f" Processing {col}...") + df[col] = df[col].apply(calculate_monthly_avg) + else: + print(f" Skipping monthly average conversion (data already monthly)") # Process list columns (expand PCT variables) print(f" Processing list columns...") @@ -703,16 +1134,20 @@ def generate_base_dataset(variable_definitions): # Save processed file df.to_pickle(file_path) - print(f" ✅ Post-processing completed for {os.path.basename(file_path)}") + print(f" Post-processing completed for {os.path.basename(file_path)}") return output_dir -def generate_base_dataset_initial_only(variable_definitions): +def generate_base_dataset_initial_only(variable_definitions, use_monthly_forcing=False, forcing_year_range="1980-1999"): """Generate base training dataset for initial-only mode (excludes Y_ variables from final_spinup files)""" print(f"\n{'='*80}") print("STEP 1: Base Dataset Generation (Initial-Only Mode)") print(f"{'='*80}") + if use_monthly_forcing: + print("📊 Using pre-computed monthly average forcing data") + print(f" Year range: {forcing_year_range}") + # File paths from config surface_data_files = config.surface_data_files ad_spinup_history_files = config.ad_spinup_history_files @@ -723,13 +1158,18 @@ def generate_base_dataset_initial_only(variable_definitions): forcing_files = {} for var_name in variable_definitions['time_series_vars']: # Look for files containing the variable name in forcing_netcdf directory - pattern = os.path.join(config.forcing_netcdf_output_dir, f'*{var_name}*1980-1999.nc') + # Support both old format (*VAR*1980-1999.nc) and new format (*VAR*2004-2023.nc) + pattern = os.path.join(config.forcing_netcdf_output_dir, f'*{var_name}*{forcing_year_range}.nc') matching_files = glob.glob(pattern) + if not matching_files: + # Try alternative pattern without year range in filename + pattern_alt = os.path.join(config.forcing_netcdf_output_dir, f'{var_name}_*.nc') + matching_files = glob.glob(pattern_alt) if matching_files: forcing_files[var_name] = matching_files[0] # Use first match - print(f"✅ Found forcing file: {os.path.basename(matching_files[0])}") + print(f"Found forcing file: {os.path.basename(matching_files[0])}") else: - print(f"⚠️ Forcing file not found for {var_name}: {pattern}") + print(f"Forcing file not found for {var_name}: {pattern}") print(f"Found {len(forcing_files)} forcing files") @@ -749,7 +1189,7 @@ def generate_base_dataset_initial_only(variable_definitions): ds_forcing = {} for var_name, file_path in forcing_files.items(): ds_forcing[var_name] = nc.Dataset(file_path) - print(f"✅ Forcing data loaded: {var_name}") + print(f"Forcing data loaded: {var_name}") print(f"All files loaded in {time.time() - start_time:.2f} seconds") @@ -786,10 +1226,10 @@ def generate_base_dataset_initial_only(variable_definitions): _, all_restart_indices = restart_tree.query(query_coords, k=1) _, all_forcing_indices = forcing_tree.query(query_coords, k=1) - print("✅ KDTree indices built") + print("KDTree indices built") # Pre-load forcing data into memory for optimization - print("🚀 Pre-loading forcing data into memory...") + print("Pre-loading forcing data into memory...") forcing_data = {} for var_name, ds in ds_forcing.items(): forcing_data[var_name] = ds.variables[var_name][:, 0, :] # (time, 1, grid_cells) @@ -810,7 +1250,7 @@ def generate_base_dataset_initial_only(variable_definitions): pft_map[grid_id] = np.where(pft_gridcell_index == grid_id)[0] column_map[grid_id] = np.where(column_gridcell_index == grid_id)[0] - print("✅ Index mappings built") + print("Index mappings built") # Process data in batches batch_size = 1000 @@ -869,13 +1309,14 @@ def generate_base_dataset_initial_only(variable_definitions): # NOTE: We do NOT add Y_ variables for initial-only mode # Process each gridcell in the batch - for k, gridcell_idx in enumerate(batch_gridcells): + for k in range(len(batch_gridcells)): if k % 100 == 0: - print(f" Processing gridcell {k}/{len(batch_gridcells)} (idx={gridcell_idx})") + print(f" Processing gridcell {k}/{len(batch_gridcells)}") # Get indices - restart_idx = batch_restart_indices[k] - forcing_idx = batch_forcing_indices[k] + restart_idx = int(batch_restart_indices[k]) # Index for ds1 (surface) and ds10 (restart) files + forcing_idx = int(batch_forcing_indices[k]) # Index for forcing files + gridcell_idx_history = batch_gridcells[k] # Index for ds2 (history) file - may be 2D flattened gridcell_id = restart_idx + 1 pft_indices_for_cell = pft_map.get(gridcell_id, []) @@ -889,63 +1330,213 @@ def generate_base_dataset_initial_only(variable_definitions): if k == 999: print(f" After 1000th gridcell: Latitude={len(data_dict['Latitude'])}, FLDS={len(data_dict.get('FLDS', []))}") - # Process forcing data (time series variables) - store raw data like original script + # Process forcing data (time series variables) for var_name in variable_definitions['time_series_vars']: if var_name in forcing_data: time_series = forcing_data[var_name][:, forcing_idx] - # Store raw time series like original script (will be processed later) - data_dict[var_name].append(time_series) + # If using monthly forcing, data is already monthly average, store as-is + # Otherwise, store raw time series (will be processed later in post-processing) + if use_monthly_forcing: + # Data is already monthly average, convert to list of floats + data_dict[var_name].append(time_series.tolist() if isinstance(time_series, np.ndarray) else list(time_series)) + else: + # Store raw time series (will be processed later) + data_dict[var_name].append(time_series) else: data_dict[var_name].append([]) # Add empty list if variable not found # Process surface properties for var_name in variable_definitions['surface_vars']: if var_name == 'Latitude': - data_dict[var_name].append(lats[gridcell_idx]) + # Use global index (start_idx + k) since lats array is already filtered for valid gridcells + global_idx = start_idx + k + data_dict[var_name].append(lats[global_idx]) elif var_name == 'Longitude': - data_dict[var_name].append(lons[gridcell_idx]) + # Use global index (start_idx + k) since lons array is already filtered for valid gridcells + global_idx = start_idx + k + data_dict[var_name].append(lons[global_idx]) elif var_name == 'landfrac': # landfrac comes from history file (ds2), not surface file (ds1) # Convert to float64 to match reference file - landfrac_val = ds2.variables['landfrac'][gridcell_idx] + global_idx = start_idx + k + if grid_format == '1d': + landfrac_val = ds2.variables['landfrac'][gridcell_idx_history] + elif grid_format == '2d': + lat_idx = grid_info['lat_idx_map'].get(global_idx, None) + lon_idx = grid_info['lon_idx_map'].get(global_idx, None) + if lat_idx is not None and lon_idx is not None: + landfrac_val = ds2.variables['landfrac'][lat_idx, lon_idx] + else: + # Fallback: use flat index + landfrac_val = ds2.variables['landfrac'].flatten()[gridcell_idx_history] + else: + landfrac_val = ds2.variables['landfrac'][gridcell_idx_history] + if hasattr(landfrac_val, 'data'): # MaskedArray landfrac_val = landfrac_val.data + # Ensure scalar value + if isinstance(landfrac_val, np.ndarray) and landfrac_val.size > 1: + landfrac_val = landfrac_val.item() if landfrac_val.size == 1 else landfrac_val[0] data_dict[var_name].append(float(landfrac_val)) elif var_name == 'PCT_CLAY': # Store PCT_CLAY as a list (all levels) - pct_clay_data = ds1.variables['PCT_CLAY'][:, gridcell_idx] - # Convert MaskedArray to regular array and ensure float64 - if hasattr(pct_clay_data, 'data'): # MaskedArray - pct_clay_data = pct_clay_data.data - data_dict[var_name].append(pct_clay_data.astype(np.float64).tolist()) + # PCT_CLAY is from ds1, check ds1's grid format and variable dimensions + var_obj = ds1.variables['PCT_CLAY'] + var_dims = var_obj.dimensions + grid_format_ds1 = grid_info['format_ds1'] + dim1_ds1 = grid_info['dim1_ds1'] + dim2_ds1 = grid_info['dim2_ds1'] + + try: + if grid_format_ds1 == '2d' and len(var_dims) == 3 and var_dims[1] == dim1_ds1 and var_dims[2] == dim2_ds1: + # Variable is (level, lat, lon) in ds1 - use 2D indexing + lat_idx = grid_info['lat_idx_map_ds1'].get(restart_idx, None) + lon_idx = grid_info['lon_idx_map_ds1'].get(restart_idx, None) + if lat_idx is not None and lon_idx is not None: + pct_clay_data = var_obj[:, lat_idx, lon_idx] + else: + # Fallback: use 1D indexing if mapping not available + pct_clay_data = var_obj[:, restart_idx] if restart_idx < var_obj.shape[-1] else var_obj[:, 0] + else: + # Variable is (level, gridcell) or 1D format - use 1D indexing with restart_idx + if len(var_dims) >= 2 and restart_idx < var_obj.shape[-1]: + pct_clay_data = var_obj[:, restart_idx] + else: + print(f" Warning: restart_idx {restart_idx} out of bounds for PCT_CLAY (shape: {var_obj.shape})") + pct_clay_data = np.zeros(var_obj.shape[0]) + + # Convert to numpy array + pct_clay_data = np.asarray(pct_clay_data) + # Convert MaskedArray to regular array and ensure float64 + if hasattr(pct_clay_data, 'data'): # MaskedArray + pct_clay_data = pct_clay_data.data + data_dict[var_name].append(pct_clay_data.astype(np.float64).tolist()) + except (IndexError, ValueError, TypeError) as e: + print(f" Error processing PCT_CLAY: {e}") + data_dict[var_name].append([]) elif var_name == 'PCT_SAND': # Store PCT_SAND as a list (all levels) - pct_sand_data = ds1.variables['PCT_SAND'][:, gridcell_idx] - # Convert MaskedArray to regular array and ensure float64 - if hasattr(pct_sand_data, 'data'): # MaskedArray - pct_sand_data = pct_sand_data.data - data_dict[var_name].append(pct_sand_data.astype(np.float64).tolist()) + # PCT_SAND is from ds1, check ds1's grid format and variable dimensions + var_obj = ds1.variables['PCT_SAND'] + var_dims = var_obj.dimensions + grid_format_ds1 = grid_info['format_ds1'] + dim1_ds1 = grid_info['dim1_ds1'] + dim2_ds1 = grid_info['dim2_ds1'] + + try: + if grid_format_ds1 == '2d' and len(var_dims) == 3 and var_dims[1] == dim1_ds1 and var_dims[2] == dim2_ds1: + # Variable is (level, lat, lon) in ds1 - use 2D indexing + lat_idx = grid_info['lat_idx_map_ds1'].get(restart_idx, None) + lon_idx = grid_info['lon_idx_map_ds1'].get(restart_idx, None) + if lat_idx is not None and lon_idx is not None: + pct_sand_data = var_obj[:, lat_idx, lon_idx] + else: + # Fallback: use 1D indexing if mapping not available + pct_sand_data = var_obj[:, restart_idx] if restart_idx < var_obj.shape[-1] else var_obj[:, 0] + else: + # Variable is (level, gridcell) or 1D format - use 1D indexing with restart_idx + if len(var_dims) >= 2 and restart_idx < var_obj.shape[-1]: + pct_sand_data = var_obj[:, restart_idx] + else: + print(f" Warning: restart_idx {restart_idx} out of bounds for PCT_SAND (shape: {var_obj.shape})") + pct_sand_data = np.zeros(var_obj.shape[0]) + + # Convert to numpy array + pct_sand_data = np.asarray(pct_sand_data) + # Convert MaskedArray to regular array and ensure float64 + if hasattr(pct_sand_data, 'data'): # MaskedArray + pct_sand_data = pct_sand_data.data + data_dict[var_name].append(pct_sand_data.astype(np.float64).tolist()) + except (IndexError, ValueError, TypeError) as e: + print(f" Error processing PCT_SAND: {e}") + data_dict[var_name].append([]) elif var_name.startswith('PCT_NAT_PFT_') or var_name.startswith('PCT_CLAY_') or var_name.startswith('PCT_SAND_'): # Handle 2D variables with level indices if '_' in var_name: level_idx = int(var_name.split('_')[-1]) base_var = '_'.join(var_name.split('_')[:-1]) # e.g., 'PCT_CLAY' if base_var in ds1.variables: - pct_val = ds1.variables[base_var][level_idx, gridcell_idx] - # Convert MaskedArray to regular array and ensure float64 - if hasattr(pct_val, 'data'): # MaskedArray - pct_val = pct_val.data - data_dict[var_name].append(float(pct_val)) + var_obj = ds1.variables[base_var] + var_dims = var_obj.dimensions + var_shape = var_obj.shape + grid_format_ds1 = grid_info['format_ds1'] + dim1_ds1 = grid_info['dim1_ds1'] + dim2_ds1 = grid_info['dim2_ds1'] + + try: + if grid_format_ds1 == '2d' and len(var_dims) == 3 and var_dims[1] == dim1_ds1 and var_dims[2] == dim2_ds1: + # Variable is (level, lat, lon) in ds1 - use 2D indexing + lat_idx = grid_info['lat_idx_map_ds1'].get(restart_idx, None) + lon_idx = grid_info['lon_idx_map_ds1'].get(restart_idx, None) + if lat_idx is not None and lon_idx is not None and level_idx < var_shape[0]: + pct_val_raw = var_obj[level_idx, lat_idx, lon_idx] + pct_val = np.asarray(pct_val_raw).item() + else: + # Fallback: use 1D indexing if mapping not available + if level_idx < var_shape[0] and restart_idx < var_shape[1]: + pct_val_raw = var_obj[level_idx, restart_idx] + pct_val = np.asarray(pct_val_raw).item() + else: + pct_val = 0.0 + else: + # Variable is (level, gridcell) or 1D format - use 1D indexing with restart_idx + if level_idx < var_shape[0] and restart_idx < var_shape[1]: + pct_val_raw = var_obj[level_idx, restart_idx] + pct_val = np.asarray(pct_val_raw).item() + else: + print(f" Warning: Index out of bounds for {var_name} (level_idx={level_idx}, restart_idx={restart_idx}, shape={var_shape})") + pct_val = 0.0 + + # Final check: ensure we have a numeric value + if not isinstance(pct_val, (int, float, np.integer, np.floating)): + pct_val = float(pct_val) + + data_dict[var_name].append(float(pct_val)) + except (IndexError, ValueError, TypeError, AttributeError) as e: + print(f" Error processing {var_name}: {e}") + data_dict[var_name].append(0.0) else: data_dict[var_name].append(0.0) else: data_dict[var_name].append(0.0) elif var_name in ds1.variables: - # 1D variables - val = ds1.variables[var_name][gridcell_idx] - # Convert MaskedArray to regular array - if hasattr(val, 'data'): # MaskedArray - val = val.data + # ds1 variables are indexed by restart_idx (gridcell index), not gridcell_idx_history + var_obj = ds1.variables[var_name] + var_dims = var_obj.dimensions + var_shape = var_obj.shape + + try: + if len(var_dims) == 1: + # Variable is 1D (gridcell) - typical for surface scalar variables + if restart_idx < var_shape[0]: + # Read value and immediately convert to scalar using np.asarray().item() + # This handles memoryview, MaskedArray, and other netCDF types + val = np.asarray(var_obj[restart_idx]).item() + else: + print(f" Error: restart_idx {restart_idx} out of bounds for {var_name} (shape: {var_shape})") + val = 0.0 + else: + # Multi-dimensional variable - surface variables should typically be 1D + # This might be incorrectly classified, but try to handle it + print(f" Warning: {var_name} has {len(var_dims)} dimensions but is in surface_vars") + # For multi-dim, assume last dimension is gridcell, take first element of other dims + if restart_idx < var_shape[-1]: + indices = [0] * (len(var_shape) - 1) + [restart_idx] + val = np.asarray(var_obj[tuple(indices)]).item() + else: + val = 0.0 + except (IndexError, ValueError, TypeError, AttributeError) as e: + print(f" Error indexing {var_name} (dims: {var_dims}, shape: {var_shape}, restart_idx: {restart_idx}): {e}") + val = 0.0 + + # Final check: ensure we have a numeric value + if not isinstance(val, (int, float, np.integer, np.floating)): + try: + val = float(val) + except (ValueError, TypeError) as e: + print(f" Error: Could not convert {var_name} to float: {val}, type: {type(val)}, error: {e}") + val = 0.0 + # Handle integer variables if var_name in ['SOIL_COLOR', 'SOIL_ORDER']: data_dict[var_name].append(int(val)) @@ -957,10 +1548,9 @@ def generate_base_dataset_initial_only(variable_definitions): # Process scalar variables from history file (X only, no Y_ variables for initial-only mode) for var_name in variable_definitions['scalar_vars']: if var_name in ds2.variables: - val = ds2.variables[var_name][0, gridcell_idx] - # Convert MaskedArray to regular array - if hasattr(val, 'data'): # MaskedArray - val = val.data + # ds2 variables are indexed by gridcell_idx_history + val_raw = ds2.variables[var_name][0, gridcell_idx_history] + val = np.asarray(val_raw).item() # Handle integer variables if var_name in ['SOIL_COLOR', 'SOIL_ORDER']: data_dict[var_name].append(int(val)) @@ -1033,7 +1623,7 @@ def generate_base_dataset_initial_only(variable_definitions): lengths = {key: len(values) for key, values in data_dict.items()} unique_lengths = set(lengths.values()) if len(unique_lengths) > 1: - print(f" ❌ Length mismatch detected!") + print(f"Length mismatch detected!") for length in unique_lengths: vars_with_length = [k for k, v in lengths.items() if v == length] print(f" Length {length}: {len(vars_with_length)} variables") @@ -1048,7 +1638,7 @@ def generate_base_dataset_initial_only(variable_definitions): batch_files.append(batch_save_path) # Add to list for post-processing batch_time = time.time() - batch_start_time - print(f"✅ Batch {batch_number} completed: {batch_time:.2f}s") + print(f" Batch {batch_number} completed: {batch_time:.2f}s") print(f" Path: {batch_save_path}") print(f" Shape: {df_batch.shape}") print(f" Columns: {len(df_batch.columns)}") @@ -1061,14 +1651,17 @@ def generate_base_dataset_initial_only(variable_definitions): ds2.close() ds10.close() - print("✅ All NetCDF files closed") - print(f"✅ Base dataset generation completed!") + print("All NetCDF files closed") + print(f"Base dataset generation completed!") print(f"Total batches: {batch_number - 1}") print(f"Output directory: {output_dir}") # Post-processing like original script print(f"\n{'='*80}") - print("POST-PROCESSING: Converting to monthly averages and expanding variables") + if use_monthly_forcing: + print("POST-PROCESSING: Expanding variables (skipping monthly average conversion - data already monthly)") + else: + print("POST-PROCESSING: Converting to monthly averages and expanding variables") print(f"{'='*80}") # Process all generated files @@ -1078,12 +1671,15 @@ def generate_base_dataset_initial_only(variable_definitions): # Load the file df = pd.read_pickle(file_path) - # Process time series columns (convert to monthly averages) - print(f" Processing time series columns...") - for col in variable_definitions['time_series_vars']: - if col in df.columns: - print(f" Processing {col}...") - df[col] = df[col].apply(calculate_monthly_avg) + # Process time series columns (convert to monthly averages only if not using pre-computed monthly data) + if not use_monthly_forcing: + print(f" Processing time series columns (converting to monthly averages)...") + for col in variable_definitions['time_series_vars']: + if col in df.columns: + print(f" Processing {col}...") + df[col] = df[col].apply(calculate_monthly_avg) + else: + print(f" Skipping monthly average conversion (data already monthly)") # Process list columns (expand PCT variables) print(f" Processing list columns...") @@ -1099,7 +1695,7 @@ def generate_base_dataset_initial_only(variable_definitions): # Save processed file df.to_pickle(file_path) - print(f" ✅ Post-processing completed for {os.path.basename(file_path)}") + print(f"Post-processing completed for {os.path.basename(file_path)}") return output_dir @@ -1126,11 +1722,11 @@ def generate_enhanced_dataset(base_output_dir, variable_definitions, initial_onl print(f"Found {len(base_files)} base PKL files to enhance") if not base_files: - print("❌ No base PKL files found") + print("No base PKL files found") return base_output_dir if initial_only_mode and (not config.final_spinup_history_files or not config.final_spinup_restart_files): - print("⚠️ Initial-only mode detected with no final spinup files; skipping enhanced dataset generation.") + print("Initial-only mode detected with no final spinup files; skipping enhanced dataset generation.") return base_output_dir # Load restart files for enhancement @@ -1229,11 +1825,11 @@ def generate_enhanced_dataset(base_output_dir, variable_definitions, initial_onl # Save enhanced dataset enhanced_file = os.path.join(enhanced_output_dir, f"enhanced_monthly_training_data_batch_{i:02d}.pkl") df_enhanced.to_pickle(enhanced_file) - print(f" ✅ Enhanced dataset saved: {os.path.basename(enhanced_file)}") - print(f" 📐 Enhanced shape: {df_enhanced.shape}") + print(f"Enhanced dataset saved: {os.path.basename(enhanced_file)}") + print(f"Enhanced shape: {df_enhanced.shape}") except Exception as e: - print(f" ❌ Failed to process file: {e}") + print(f"Failed to process file: {e}") continue return enhanced_output_dir @@ -1257,7 +1853,7 @@ def add_pft_variables(enhanced_output_dir, variable_definitions): print(f"File path: {config.clm_params_nc_path}") if not os.path.exists(config.clm_params_nc_path): - print(f"❌ Error: CLM parameters file not found: {config.clm_params_nc_path}") + print(f"Error: CLM parameters file not found: {config.clm_params_nc_path}") return enhanced_output_dir ds = nc.Dataset(config.clm_params_nc_path) @@ -1290,7 +1886,7 @@ def add_pft_variables(enhanced_output_dir, variable_definitions): else: print(f"Skipped {var}: not found in NetCDF") - print(f"\n✅ Successfully loaded {len(broadcast_feature_dict)} PFT variables from NetCDF") + print(f"\nSuccessfully loaded {len(broadcast_feature_dict)} PFT variables from NetCDF") # Get enhanced PKL files - handle different naming patterns input_files = sorted(glob.glob(os.path.join(enhanced_output_dir, "enhanced_monthly_training_data_batch_*.pkl"))) @@ -1317,7 +1913,7 @@ def add_pft_variables(enhanced_output_dir, variable_definitions): # Check if PFT variables already exist existing_pft_cols = [col for col in df.columns if col.startswith("pft_")] if existing_pft_cols: - print(f" ⚠️ File already contains {len(existing_pft_cols)} PFT variables, skipping addition") + print(f"File already contains {len(existing_pft_cols)} PFT variables, skipping addition") continue # Add each variable as a vector column with pft_ prefix @@ -1326,20 +1922,20 @@ def add_pft_variables(enhanced_output_dir, variable_definitions): df["pft_" + var] = [val_list] * len(df) # Add the same list to each row new_shape = df.shape - print(f" ✅ Successfully added {len(broadcast_feature_dict)} PFT variables") - print(f" 📐 New data shape: {original_shape} → {new_shape}") + print(f"Successfully added {len(broadcast_feature_dict)} PFT variables") + print(f"New data shape: {original_shape} → {new_shape}") # Save in-place (overwrite original file) df.to_pickle(file_path) - print(f" ✅ File saved: {os.path.basename(file_path)}") + print(f"File saved: {os.path.basename(file_path)}") except Exception as e: - print(f" ❌ Failed to process file: {e}") + print(f"Failed to process file: {e}") continue ds.close() - print(f"\n✅ PFT variables addition completed!") + print(f"\nPFT variables addition completed!") print(f" - Total files: {len(input_files)}") print(f" - PFT variables added: {len(broadcast_feature_dict)}") @@ -1386,7 +1982,7 @@ def final_variable_cleanup(enhanced_output_dir, variable_definitions): print(f"\n🔍 Found {len(input_files)} enhanced PKL files to process") if len(input_files) == 0: - print("❌ No enhanced PKL files found.") + print("No enhanced PKL files found.") return enhanced_output_dir # Process each PKL file @@ -1428,24 +2024,24 @@ def final_variable_cleanup(enhanced_output_dir, variable_definitions): df_cleaned = df[list(vars_to_keep)] new_shape = df_cleaned.shape - print(f" ✅ Successfully removed {len(vars_to_remove)} variables") - print(f" 📐 New data shape: {original_shape} → {new_shape}") + print(f"Successfully removed {len(vars_to_remove)} variables") + print(f"New data shape: {original_shape} → {new_shape}") # Save in-place (overwrite original file) df_cleaned.to_pickle(file_path) - print(f" ✅ File saved: {os.path.basename(file_path)}") + print(f"File saved: {os.path.basename(file_path)}") else: - print(" ℹ️ No extra variables found to remove. File already matches CNP_IO file.") + print("No extra variables found to remove. File already matches CNP_IO file.") if len(current_vars) != len(all_expected_vars): - print(f" ⚠️ Warning: Column count mismatch. Current: {len(current_vars)}, Expected: {len(all_expected_vars)}") + print(f" Warning: Column count mismatch. Current: {len(current_vars)}, Expected: {len(all_expected_vars)}") print(f" Missing from current: {list(all_expected_vars - current_vars)[:5]}...") print(f" Extra in current: {list(current_vars - all_expected_vars)[:5]}...") except Exception as e: - print(f" ❌ Failed to process file: {e}") + print(f"Failed to process file: {e}") continue - print(f"\n✅ Variable cleanup completed!") + print(f"\nVariable cleanup completed!") print(f" - Total files: {len(input_files)}") print(f" - Kept only variables defined in CNP_IO file") print(f" - Target variable count: {len(all_expected_vars)}") @@ -1472,9 +2068,9 @@ def generate_forcing_only_dataset(): matching_files = glob.glob(pattern) if matching_files: forcing_data_files[var_name] = matching_files[0] # Use first match - print(f"✅ Found forcing file: {os.path.basename(matching_files[0])}") + print(f"Found forcing file: {os.path.basename(matching_files[0])}") else: - print(f"⚠️ Forcing file not found for {var_name}: {pattern}") + print(f"Forcing file not found for {var_name}: {pattern}") print(f"Found {len(forcing_data_files)} forcing files") @@ -1522,10 +2118,10 @@ def generate_forcing_only_dataset(): forcing_tree = cKDTree(forcing_coords) _, all_forcing_indices = forcing_tree.query(query_coords, k=1) - print("✅ KDTree indices built") + print("KDTree indices built") # Pre-load forcing data into memory for optimization - print("🚀 Pre-loading forcing data into memory...") + print("Pre-loading forcing data into memory...") forcing_data = {} for var_name, ds in ds_forcing.items(): forcing_data[var_name] = ds.variables[var_name][:, 0, :] # (time, 1, grid_cells) @@ -1598,7 +2194,7 @@ def generate_forcing_only_dataset(): batch_files.append(batch_save_path) # Add to list for tracking batch_time = time.time() - batch_start_time - print(f"✅ Batch {batch_number} completed: {batch_time:.2f}s") + print(f" Batch {batch_number} completed: {batch_time:.2f}s") print(f" Path: {batch_save_path}") print(f" Shape: {df_batch.shape}") print(f" Forcing data length: {len(df_batch['FLDS'].iloc[0])}") @@ -1610,8 +2206,8 @@ def generate_forcing_only_dataset(): ds1.close() ds2.close() - print("✅ All NetCDF files closed") - print(f"✅ Forcing-only dataset generation completed!") + print("All NetCDF files closed") + print(f"Forcing-only dataset generation completed!") print(f"Total batches: {batch_number - 1}") print(f"Output directory: {output_dir}") @@ -1636,6 +2232,17 @@ def main(): action="store_true", help="Generate initial condition dataset (excludes Y_ variables from final_spinup files)" ) + parser.add_argument( + "--use_monthly_forcing", + action="store_true", + help="Use pre-computed monthly average forcing data (skip calculate_monthly_avg processing)" + ) + parser.add_argument( + "--forcing_year_range", + type=str, + default="1980-1999", + help="Year range for forcing files (e.g., '1980-1999' or '2004-2023'). Default: 1980-1999" + ) args = parser.parse_args() # Validate arguments @@ -1682,7 +2289,11 @@ def main(): variable_definitions = parse_cnp_io_variables() # Step 1: Generate base dataset - base_output_dir = generate_base_dataset(variable_definitions) + base_output_dir = generate_base_dataset( + variable_definitions, + use_monthly_forcing=args.use_monthly_forcing, + forcing_year_range=args.forcing_year_range + ) # Step 2: Generate enhanced dataset enhanced_output_dir = generate_enhanced_dataset(base_output_dir, variable_definitions) @@ -1720,7 +2331,11 @@ def main(): variable_definitions = parse_cnp_io_variables() # Step 1: Generate base dataset (without Y_ variables) - base_output_dir = generate_base_dataset_initial_only(variable_definitions) + base_output_dir = generate_base_dataset_initial_only( + variable_definitions, + use_monthly_forcing=args.use_monthly_forcing, + forcing_year_range=args.forcing_year_range + ) # Step 2: Generate enhanced dataset (same as regular enhanced dataset) enhanced_output_dir = generate_enhanced_dataset(base_output_dir, variable_definitions, initial_only_mode=True)