from typing import Any from logging import Logger import pandas as pd from src.scale_processor import ScaleProcessor from src.composite_processor import process_composites from src.utils.data_loader import assemble_wave_info, load_yaml class DataPreprocessingAllWaves: """Class for preprocessing data across all waves of the study. This class loads data and configuration for each wave, processes scales and composites, and returns preprocessed DataFrames for each wave. """ def __init__( self, data_with_configs: dict, settings: dict[str, Any], logger: Logger ): """Initialize the preprocessing class with data and settings. Args: data_with_configs (dict): Dictionary mapping wave numbers to their data and config paths. settings (dict[str, Any]): Project settings loaded from the settings file. """ self.data_with_configs: dict = data_with_configs self.settings: dict[str, Any] = settings self.logger: Logger = logger self.cronbachs_alphas: dict[str, dict[int, float]] = {} def _aggregate_cronbachs_alpha_values( self, scale_name: str, alpha_value: float | None, wave_number: int, coalesced: bool = False, ) -> None: """Aggregate Cronbach's alpha values across waves. Args: scale_name (str): Name of the scale. alpha_value (float | None): Cronbach's alpha value for the scale. wave_number (int): Current wave number. coalesced (bool): Whether this is a coalesced composite scale. """ if alpha_value is None: return if scale_name not in self.cronbachs_alphas: self.cronbachs_alphas[scale_name] = {} self.cronbachs_alphas[scale_name][wave_number] = alpha_value def preprocess_data(self) -> dict[int, pd.DataFrame]: """Preprocess data for all waves. Loads configuration for each wave, processes scales and composite scales, and returns a dictionary of preprocessed DataFrames indexed by wave number. Returns: dict[int, pd.DataFrame]: Dictionary mapping wave numbers to their preprocessed DataFrames. Raises: ValueError: If required configuration keys or columns are missing. """ all_preprocessed: dict = {} for wave_number, data_of_wave in self.data_with_configs.items(): data = data_of_wave["data"] config_path = data_of_wave["config_path"] wave_config = load_yaml(config_path) participant_id_column = wave_config.get("participant_id_column") if participant_id_column is None: raise ValueError( f"Wave {wave_number}: Required key 'participant_id_column' missing in config '{config_path}'." ) if participant_id_column not in data.columns: raise ValueError( f"Wave {wave_number}: Participant ID column '{participant_id_column}' not found in the data for config '{config_path}'." ) ( scale_dict, subgroup_scales, skip_scales, composite_scales, ) = assemble_wave_info(config_path, self.settings) scale_dfs: list = [] all_scale_outputs: list = [] scale_item_counts: dict[str, int] = {} for scale_name, subgroup in subgroup_scales.items(): if scale_name in skip_scales: continue if scale_name not in scale_dict: raise ValueError( f"Scale {scale_name} not in loaded scale configs (check YAML)." ) scale_config = scale_dict[scale_name] number_items = len(scale_config.get("items", [])) output_scale_name = scale_config.get("output", scale_name) scale_item_counts[output_scale_name] = number_items scale_processor: ScaleProcessor = ScaleProcessor( scale_config, logger=self.logger, subgroup_name=subgroup ) scale_dataframe: pd.DataFrame = scale_processor.process(data) scale_dfs.append(scale_dataframe) all_scale_outputs.extend(scale_dataframe.columns.tolist()) output_name = scale_processor.output self._aggregate_cronbachs_alpha_values( output_name, scale_processor.cronbachs_alpha, wave_number, coalesced=False, ) for part_name, alpha_value in getattr( scale_processor, "cronbachs_alpha_by_part", {} ).items(): subscale_column = f"{scale_processor.name}_{part_name}_mean" self._aggregate_cronbachs_alpha_values( subscale_column, alpha_value, wave_number, coalesced=False, ) result_dataframe: pd.DataFrame = pd.concat( [data[[participant_id_column]], *scale_dfs], axis=1 ) constituent_outputs: set = set() if composite_scales: wave_alpha_dict = { scale_name: waves.get(wave_number) for scale_name, waves in self.cronbachs_alphas.items() if wave_number in waves } composite_dataframe, updated_alphas = process_composites( result_dataframe, composite_scales, wave_alpha_dict, scale_item_counts, ) for scale_name, alpha_value in updated_alphas.items(): self._aggregate_cronbachs_alpha_values( scale_name, alpha_value, wave_number, coalesced=True ) composite_output_names: list = list(composite_dataframe.columns) for composite_scale in composite_scales.values(): if composite_scale.get("keep_subscales", False): continue if "scales" in composite_scale: constituent_outputs.update(composite_scale["scales"]) result_dataframe = pd.concat( [result_dataframe, composite_dataframe], axis=1 ) columns_to_keep: list = ( [participant_id_column] + composite_output_names + [ col for col in result_dataframe.columns if col not in constituent_outputs and col not in composite_output_names and col != participant_id_column ] ) result_dataframe = result_dataframe.loc[:, columns_to_keep] all_preprocessed[wave_number] = result_dataframe return all_preprocessed