Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
21 changes: 13 additions & 8 deletions config/training_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand All @@ -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,
Expand Down Expand Up @@ -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(
Expand All @@ -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
Expand Down
63 changes: 46 additions & 17 deletions data/data_loader_individual.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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,
}
Expand All @@ -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)]
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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:
Expand All @@ -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']
Expand Down
52 changes: 32 additions & 20 deletions models/cnp_combined_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
"""
Expand All @@ -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)
Expand All @@ -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))
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(),
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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()
Expand All @@ -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
Expand Down
13 changes: 13 additions & 0 deletions train_cnp_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down Expand Up @@ -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(
Expand All @@ -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
)
Expand Down Expand Up @@ -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
)

Expand Down Expand Up @@ -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,
Expand Down
Loading
Loading