From c9ae482174132cfffd9ffa67688e7f48a86330b6 Mon Sep 17 00:00:00 2001 From: Erick Ahmed Date: Mon, 13 Jul 2026 21:53:20 +0200 Subject: [PATCH] Add positional argument to choose correlation methods (Pearson or Spearman) --- stat/correlation.py | 25 +++++++++++++++++++------ 1 file changed, 19 insertions(+), 6 deletions(-) diff --git a/stat/correlation.py b/stat/correlation.py index d492534..dc81132 100644 --- a/stat/correlation.py +++ b/stat/correlation.py @@ -10,9 +10,21 @@ import pandas as pd import plotly.graph_objects as go from utils.extractor import load_data -from utils.extractor import to_int -def calculate_correlation(df: pd.DataFrame, target_id: str | None = None) -> pd.DataFrame: + +def _to_int(x): + """Convert a hex string or integer to int, returning NaN on failure.""" + if isinstance(x, (int, np.integer)): + return int(x) + if isinstance(x, str): + try: + return int(x, 16) + except ValueError: + return np.nan + return np.nan + + +def calculate_correlation(df: pd.DataFrame, method: str, target_id: str | None = None) -> pd.DataFrame: """Calculates inter-byte correlation grouped by identifier.""" byte_cols = [f"b{i}" for i in range(8)] available_cols = [col for col in byte_cols if col in df.columns] @@ -22,16 +34,16 @@ def calculate_correlation(df: pd.DataFrame, target_id: str | None = None) -> pd. df_bytes = df[["Identifier"] + available_cols].copy() for col in available_cols: - df_bytes[col] = df_bytes[col].apply(to_int) + df_bytes[col] = df_bytes[col].apply(_to_int) if target_id: group = df_bytes[df_bytes["Identifier"] == target_id] if group.empty: raise ValueError(f"Identifier '{target_id}' not found in data") - return group[available_cols].corr().fillna(0.0) + return group[available_cols].corr(method=method).fillna(0.0) def max_abs_corr(group: pd.DataFrame) -> pd.Series: - corr_arr = np.abs(group.corr().to_numpy().copy()) + corr_arr = np.abs(group.corr(method=method).to_numpy().copy()) np.fill_diagonal(corr_arr, 0.0) return pd.Series(corr_arr.max(axis=0), index=group.columns).fillna(0.0) @@ -109,6 +121,7 @@ def plot_correlation_heatmap(corr_df: pd.DataFrame, target_id: str | None, title if __name__ == "__main__": parser = argparse.ArgumentParser(description="Analyze CAN bus inter-byte correlation") + parser.add_argument("method", choices=["pearson", "spearman"], help="Correlation method to use") parser.add_argument("input", type=Path, help="Path to the input CAN log file") parser.add_argument("output", type=Path, nargs="?", default=Path("correlation_report.html"), help="Path to the output HTML report") parser.add_argument("title", nargs="?", default="CAN Bus Inter-Byte Correlation", help="Title for the HTML report") @@ -116,7 +129,7 @@ if __name__ == "__main__": args = parser.parse_args() df = load_data(args.input) - corr_df = calculate_correlation(df, target_id=args.identifier) + corr_df = calculate_correlation(df, method=args.method, target_id=args.identifier) display_title = f"{args.title} ({args.identifier})" if args.identifier else args.title fig = plot_correlation_heatmap(corr_df, target_id=args.identifier, title=display_title)