""" EDA Module for LLM-Enhanced Predictive Maintenance Author: Antigravity AI Date: August 2026 This module generates exploratory data analysis plots and reports, saving them to disk: 1. Missing value analysis (heatmaps and counts). 2. Class distribution bar charts. 3. Feature-to-target correlation heatmaps. 4. Pre-modeling statistical feature relevance (using Mutual Information). 5. Failure distribution timelines. 6. Sensor behavior overlay plots (pre-failure window vs. normal operation). """ import os import time import logging import matplotlib.pyplot as plt import seaborn as sns import numpy as np import pandas as pd from sklearn.feature_selection import mutual_info_classif try: from utils.project_paths import LOGS_DIR except ImportError: from project_paths import LOGS_DIR # Setup logging os.makedirs(LOGS_DIR, exist_ok=True) logging.basicConfig( level=logging.INFO, format='%(asctime)s [%(levelname)s] %(message)s', handlers=[ logging.StreamHandler(), logging.FileHandler(os.path.join(LOGS_DIR, "eda.log")) ] ) logger = logging.getLogger(__name__) def time_tracker(func): """Decorator to measure and log the execution time of functions.""" def wrapper(*args, **kwargs): start_time = time.time() logger.info(f"Starting execution of function: {func.__name__}") result = func(*args, **kwargs) elapsed_time = time.time() - start_time logger.info(f"Finished execution of: {func.__name__}. Time taken: {elapsed_time:.2f} seconds.") return result return wrapper @time_tracker def plot_missing_values(df: pd.DataFrame, dataset_name: str, output_dir: str): """Analyze and generate missing value plots.""" os.makedirs(output_dir, exist_ok=True) logger.info(f"Analyzing missing values for {dataset_name}...") missing_count = df.isnull().sum() missing_percent = (missing_count / len(df)) * 100 missing_df = pd.DataFrame({'Missing Count': missing_count, 'Percentage': missing_percent}) missing_df = missing_df[missing_df['Missing Count'] > 0].sort_values(by='Missing Count', ascending=False) if missing_df.empty: logger.info(f"No missing values found in {dataset_name}.") return plt.figure(figsize=(12, 6)) sns.barplot(x=missing_df.index[:30], y=missing_df['Percentage'][:30], palette='viridis') plt.xticks(rotation=90) plt.title(f"Missing Value Percentages in {dataset_name} (Top 30)") plt.ylabel("Percentage (%)") plt.tight_layout() save_path = os.path.join(output_dir, f"{dataset_name}_missing_values.png") plt.savefig(save_path, dpi=150) plt.close() logger.info(f"Saved missing value analysis to {save_path}") @time_tracker def plot_class_distributions(df: pd.DataFrame, target_col: str, dataset_name: str, output_dir: str): """Generate class distribution bar charts for classification datasets.""" os.makedirs(output_dir, exist_ok=True) logger.info(f"Plotting class distribution for {dataset_name} target '{target_col}'...") class_counts = df[target_col].value_counts() class_percents = df[target_col].value_counts(normalize=True) * 100 plt.figure(figsize=(8, 5)) ax = sns.countplot(x=target_col, data=df, palette='Set2') plt.title(f"Class Distribution: {dataset_name}") plt.ylabel("Count") plt.xlabel("Class") # Annotate percentages for i, p in enumerate(ax.patches): height = p.get_height() ax.text(p.get_x() + p.get_width()/2., height + 0.01*len(df), f'{class_counts.iloc[i]} ({class_percents.iloc[i]:.2f}%)', ha="center") plt.tight_layout() save_path = os.path.join(output_dir, f"{dataset_name}_class_distribution.png") plt.savefig(save_path, dpi=150) plt.close() logger.info(f"Saved class distribution plot to {save_path}") @time_tracker def plot_correlation_heatmap(df: pd.DataFrame, feature_cols: list, dataset_name: str, output_dir: str): """Plot correlation matrix heatmap for numeric features.""" os.makedirs(output_dir, exist_ok=True) logger.info(f"Generating correlation heatmap for {dataset_name}...") # Check shape to prevent huge slow heatmap plots for broad datasets if len(feature_cols) > 30: logger.info("Too many features for dense heatmap. Plotting top 25 features based on variance.") top_variance_cols = df[feature_cols].var().sort_values(ascending=False).index[:25].tolist() corr_matrix = df[top_variance_cols].corr() else: corr_matrix = df[feature_cols].corr() plt.figure(figsize=(14, 10)) sns.heatmap(corr_matrix, annot=True if len(corr_matrix) <= 15 else False, fmt=".2f", cmap="coolwarm", cbar=True, square=True) plt.title(f"Correlation Heatmap: {dataset_name}") plt.tight_layout() save_path = os.path.join(output_dir, f"{dataset_name}_correlation_heatmap.png") plt.savefig(save_path, dpi=150) plt.close() logger.info(f"Saved correlation heatmap to {save_path}") @time_tracker def plot_statistical_feature_importance(df: pd.DataFrame, feature_cols: list, target_col: str, dataset_name: str, output_dir: str): """Compute and plot pre-model statistical feature importance using Mutual Information.""" os.makedirs(output_dir, exist_ok=True) logger.info(f"Computing Mutual Information scores for {dataset_name}...") # Drop rows with NaNs for statistical check df_clean = df[[target_col] + feature_cols].dropna() X = df_clean[feature_cols] y = df_clean[target_col] # Compute Mutual Information classification scores mi_scores = mutual_info_classif(X, y, random_state=42) mi_df = pd.DataFrame({'Feature': feature_cols, 'MI Score': mi_scores}) mi_df = mi_df.sort_values(by='MI Score', ascending=False).reset_index(drop=True) plt.figure(figsize=(10, 8)) sns.barplot(x='MI Score', y='Feature', data=mi_df.head(25), palette='magma') plt.title(f"Mutual Information Feature Importance (Top 25): {dataset_name}") plt.xlabel("Mutual Information Score") plt.tight_layout() save_path = os.path.join(output_dir, f"{dataset_name}_mutual_info_importance.png") plt.savefig(save_path, dpi=150) plt.close() logger.info(f"Saved Mutual Information plot to {save_path}") logger.info(f"Top 5 statistical features for {dataset_name}:\n{mi_df.head(5)}") @time_tracker def plot_failure_distribution_timeline(df: pd.DataFrame, timestamp_col: str, failure_col: str, dataset_name: str, output_dir: str): """Plot distribution of failure events over time series timeline.""" os.makedirs(output_dir, exist_ok=True) logger.info(f"Plotting failure timeline for {dataset_name}...") df_plot = df.copy() df_plot[timestamp_col] = pd.to_datetime(df_plot[timestamp_col]) df_plot = df_plot.sort_values(timestamp_col) plt.figure(figsize=(16, 4)) plt.plot(df_plot[timestamp_col], df_plot[failure_col], color='grey', alpha=0.5, label='Status / At-risk Window') # Highlight failure starting points failures = df_plot[df_plot[failure_col] == 1] plt.scatter(failures[timestamp_col], failures[failure_col], color='red', s=5, label='Failure / Danger State', zorder=5) plt.title(f"Failure Timeline Distribution: {dataset_name}") plt.xlabel("Time") plt.ylabel("Failure Flag") plt.legend() plt.grid(True, linestyle='--', alpha=0.6) plt.tight_layout() save_path = os.path.join(output_dir, f"{dataset_name}_failure_timeline.png") plt.savefig(save_path, dpi=150) plt.close() logger.info(f"Saved timeline plot to {save_path}") @time_tracker def plot_sensor_behavior_overlay(df: pd.DataFrame, sensor_cols: list, timestamp_col: str, failure_col: str, dataset_name: str, output_dir: str, window_hours: int = 24): """ Generate overlay plots comparing sensor readings in a window prior to a failure against normal, steady-state operations. Args: df: Raw or partially preprocessed dataframe containing sensor logs. sensor_cols: Sensor column labels to analyze. timestamp_col: Datetime column label. failure_col: Binary failure or 'machine_status' label. dataset_name: Name of dataset for output filing. output_dir: Plot export path. window_hours: Hour window preceding failure. """ os.makedirs(output_dir, exist_ok=True) logger.info(f"Generating sensor overlay behavior analysis for {dataset_name}...") df_sorted = df.copy() df_sorted[timestamp_col] = pd.to_datetime(df_sorted[timestamp_col]) df_sorted = df_sorted.sort_values(timestamp_col).reset_index(drop=True) # Find exact failure indices (transition from 0 to 1, or status == BROKEN) # Let's check status transitions if 'machine_status' in df_sorted.columns: failure_indices = df_sorted[df_sorted['machine_status'] == 'BROKEN'].index else: # Check transition to failure failure_indices = df_sorted[(df_sorted[failure_col] == 1) & (df_sorted[failure_col].shift(1) == 0)].index if len(failure_indices) == 0: logger.info(f"No failure events found to construct overlay windows for {dataset_name}.") return # Pick a few key sensors (up to 3) for plotting overlay to prevent overcrowded plots key_sensors = sensor_cols[:3] # Let's say hourly data (like Azure) or minute data (like Pump) # We want to identify length of time window in rows # Estimate time difference between consecutive rows time_diff_min = (df_sorted[timestamp_col].iloc[1] - df_sorted[timestamp_col].iloc[0]).total_seconds() / 60 rows_in_window = int((window_hours * 60) / time_diff_min) if time_diff_min > 0 else window_hours logger.info(f"Generating overlays using windows of {rows_in_window} rows (~{window_hours} hours).") # Build failure windows failure_windows = [] for idx in failure_indices: start_idx = max(0, idx - rows_in_window) if idx - start_idx == rows_in_window: window_df = df_sorted.loc[start_idx:idx-1, key_sensors].reset_index(drop=True) failure_windows.append(window_df) # Build normal windows (segments far away from any failures) normal_windows = [] normal_indices = df_sorted[df_sorted[failure_col] == 0].index # Draw sample normal contiguous sequences # A simple method: find chunks of normal indices longer than rows_in_window chunk_start = 0 for i in range(1, len(normal_indices)): if normal_indices[i] != normal_indices[i-1] + 1: chunk_len = normal_indices[i-1] - chunk_start if chunk_len >= rows_in_window: window_df = df_sorted.loc[chunk_start:chunk_start+rows_in_window-1, key_sensors].reset_index(drop=True) normal_windows.append(window_df) if len(normal_windows) >= len(failure_windows): break chunk_start = normal_indices[i] if not failure_windows or not normal_windows: logger.warning("Insufficient data segments to generate failure vs normal overlay analysis.") return # Aggregate windows (mean and standard deviation) fail_agg = pd.concat(failure_windows).groupby(level=0).mean() fail_std = pd.concat(failure_windows).groupby(level=0).std() norm_agg = pd.concat(normal_windows).groupby(level=0).mean() norm_std = pd.concat(normal_windows).groupby(level=0).std() # Generate overlay plot for each selected key sensor fig, axes = plt.subplots(len(key_sensors), 1, figsize=(14, 4 * len(key_sensors)), sharex=True) if len(key_sensors) == 1: axes = [axes] # X-axis represents time elapsed before failure (in hours) time_scale = np.linspace(-window_hours, 0, rows_in_window) for i, sensor in enumerate(key_sensors): # Normal operations plot axes[i].plot(time_scale, norm_agg[sensor], label='Normal Operation (Avg)', color='green', lw=2) axes[i].fill_between(time_scale, norm_agg[sensor] - norm_std[sensor], norm_agg[sensor] + norm_std[sensor], color='green', alpha=0.15) # Pre-failure plot axes[i].plot(time_scale, fail_agg[sensor], label='Pre-Failure Window (Avg)', color='red', lw=2, linestyle='--') axes[i].fill_between(time_scale, fail_agg[sensor] - fail_std[sensor], fail_agg[sensor] + fail_std[sensor], color='red', alpha=0.15) axes[i].set_title(f"Sensor Trends: {sensor} (Normal vs Pre-Failure Window)") axes[i].set_ylabel("Normalized Value") axes[i].grid(True, linestyle=':', alpha=0.7) axes[i].legend() axes[-1].set_xlabel("Hours Relative to Failure Point (0 = Failure Hour)") plt.tight_layout() save_path = os.path.join(output_dir, f"{dataset_name}_sensor_overlay_comparison.png") plt.savefig(save_path, dpi=150) plt.close() logger.info(f"Saved sensor overlay comparison plot to {save_path}") if __name__ == '__main__': logger.info("EDA module template loaded. Import and use reporting/plotting functions.")