SyedaArisha's picture
Upload folder using huggingface_hub
baf834b verified
Raw History Blame Contribute Delete
13.4 kB
"""
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.")