Compare commits
69
Commits
a3119e4fc3
...
main
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0702864bc3 | ||
|
|
ac2c88c339 | ||
|
|
96b3e35248 | ||
|
|
0483c3fea3 | ||
|
|
3701d11c77 | ||
|
|
145a5ecf78 | ||
|
|
eba9b07487 | ||
|
|
4df1187f84 | ||
|
|
c439e35e39 | ||
|
|
9406be3c18 | ||
|
|
f344808802 | ||
|
|
10fdafa244 | ||
|
|
314c4433d3 | ||
|
|
b252082991 | ||
|
|
2ec0af5f62 | ||
|
|
9c2619daa9 | ||
|
|
910e642398 | ||
|
|
0e6f39556b | ||
|
|
fb0d39c668 | ||
|
|
0b2c629d16 | ||
|
|
ef785283f0 | ||
|
|
182fc102de | ||
|
|
a064f6cc90 | ||
|
|
a4b7190756 | ||
|
|
f95d59e44d | ||
|
|
de12c1407c | ||
|
|
537b452449 | ||
|
|
3169c29319 | ||
|
|
13bd76631f | ||
|
|
6cc38291df | ||
|
|
8b6c547387 | ||
|
|
de0084dc09 | ||
|
|
af3f9d16b2 | ||
|
|
3d8c7c6639 | ||
|
|
7b7f89cf9d | ||
|
|
8f24adbdbd | ||
|
|
36bae270a1 | ||
|
|
4cb06d0497 | ||
|
|
f19dde3f9a | ||
|
|
e69000fbd8 | ||
|
|
7a63c7acd3 | ||
|
|
0f11a88ae7 | ||
|
|
984ef89a07 | ||
|
|
15190ac52e | ||
|
|
42965a4733 | ||
|
|
4eab3c9876 | ||
|
|
2b01085a9e | ||
|
|
0088cef32a | ||
|
|
cf88f88814 | ||
|
|
2a014e1e4e | ||
|
|
0e25ba4a3e | ||
|
|
832a765575 | ||
|
|
3d86bfe6d0 | ||
|
|
9b7bb945bc | ||
|
|
a9ff3880e2 | ||
|
|
5a216b22fd | ||
|
|
0294d4e584 | ||
|
|
4f6c3b7370 | ||
|
|
9951d8b4f9 | ||
|
|
5f2db4d0c9 | ||
|
|
eee173dc0b | ||
|
|
38e9354c42 | ||
|
|
ee648f9adc | ||
|
|
29d70ce713 | ||
|
|
d79fad909c | ||
|
|
fc6c593f6b | ||
|
|
10bc7c568a | ||
|
|
267df136dd | ||
|
|
f0affaed05 |
@@ -3,4 +3,9 @@
|
||||
!*.py
|
||||
!*.ipynb
|
||||
!*.md
|
||||
!*.parquet
|
||||
!.gitignore
|
||||
!*.service
|
||||
!*.timer
|
||||
!*.yaml
|
||||
!*.txt
|
||||
+3
-4
@@ -23,7 +23,7 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"file_path = \"adabase-public-0020-v_0_0_2.h5py\""
|
||||
"file_path = \"YOUR_FILE_PATH.h5py\""
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -87,7 +87,7 @@
|
||||
"id": "a4731c56",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Actions units"
|
||||
"Insights on actions units"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -167,7 +167,7 @@
|
||||
"id": "332740a8",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Plots"
|
||||
"Example plot of ECG curve"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -177,7 +177,6 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# df_signals_ecg = pd.read_hdf(file_path, \"SIGNALS\", mode=\"r\", columns=[\"STUDY\",\"LEVEL\", \"PHASE\", 'RAW_ECG_I'])\n",
|
||||
"df_signals_ecg = df_signals[[\"STUDY\",\"LEVEL\", \"PHASE\", 'RAW_ECG_I']]\n",
|
||||
"df_signals_ecg.shape"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,98 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "b9144326",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Calculate replacement values for live deployment"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "2a7b60d6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"from pathlib import Path\n",
|
||||
"import yaml"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "197cb8a6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# TODO: insert path to database\n",
|
||||
"dataset_path = Path(r\"\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "6cf26eb2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df=pd.read_parquet(dataset_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "f2c5679e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pd.set_option(\"display.max_rows\", None)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e88981b2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"\n",
|
||||
"# optional: your dataframe filtering code ...\n",
|
||||
"\n",
|
||||
"medians = df.median()\n",
|
||||
"median_dict = medians.to_dict()\n",
|
||||
"\n",
|
||||
"# Wrap in fallback key\n",
|
||||
"output = {'fallback': median_dict}\n",
|
||||
"\n",
|
||||
"# Save to YAML\n",
|
||||
"with open('config.yaml', 'w') as f:\n",
|
||||
" yaml.dump(output, f, default_flow_style=False, sort_keys=False)"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "310",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.10.19"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,611 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "89d81009",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Imports"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "7440a5b3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"from pathlib import Path\n",
|
||||
"from sklearn.preprocessing import StandardScaler, MinMaxScaler"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "09b7d707",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Config"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "2401aaef",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dataset_path = Path(r\"\") # TODO: enter path to dataset"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "0282b0b1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"FILTER_MAD = True\n",
|
||||
"THRESHOLD = 3.5\n",
|
||||
"METHOD = 'minmax'\n",
|
||||
"SCOPE = 'subject'\n",
|
||||
"FILTER_SUBSETS = True"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "a8f1716b",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Calculations"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "ac32444a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.read_parquet(dataset_path)\n",
|
||||
"df.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "3ba4401c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"if(FILTER_SUBSETS):\n",
|
||||
" # Special filter: Keep only specific subsets\n",
|
||||
"# - k-drive L1 baseline\n",
|
||||
"# - n-back L1 baseline \n",
|
||||
"# - k-drive test with levels 1, 2, 3\n",
|
||||
"\n",
|
||||
" df = df[\n",
|
||||
" (\n",
|
||||
" # k-drive L1 baseline\n",
|
||||
" ((df['STUDY'] == 'k-drive') & \n",
|
||||
" (df['LEVEL'] == 1) & \n",
|
||||
" (df['PHASE'] == 'baseline'))\n",
|
||||
" ) | \n",
|
||||
" (\n",
|
||||
" # n-back L1 baseline\n",
|
||||
" ((df['STUDY'] == 'n-back') & \n",
|
||||
" (df['LEVEL'] == 1) & \n",
|
||||
" (df['PHASE'] == 'baseline'))\n",
|
||||
" ) | \n",
|
||||
" (\n",
|
||||
" # k-drive test with levels 1, 2, 3\n",
|
||||
" ((df['STUDY'] == 'k-drive') & \n",
|
||||
" (df['LEVEL'].isin([1, 2, 3])) & \n",
|
||||
" (df['PHASE'] == 'test'))\n",
|
||||
" )].copy()\n",
|
||||
"\n",
|
||||
"print(f\"Filtered dataframe shape: {df.shape}\")\n",
|
||||
"print(f\"Remaining subsets: {df.groupby(['STUDY', 'LEVEL', 'PHASE']).size()}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "77dbd6df",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"face_au_cols = [c for c in df.columns if c.startswith(\"FACE_AU\")]\n",
|
||||
"eye_cols = ['Fix_count_short_66_150', 'Fix_count_medium_300_500',\n",
|
||||
" 'Fix_count_long_gt_1000', 'Fix_count_100', 'Fix_mean_duration',\n",
|
||||
" 'Fix_median_duration', 'Sac_count', 'Sac_mean_amp', 'Sac_mean_dur',\n",
|
||||
" 'Sac_median_dur', 'Blink_count', 'Blink_mean_dur', 'Blink_median_dur',\n",
|
||||
" 'Pupil_mean', 'Pupil_IPA']\n",
|
||||
"eye_cols_without_blink = ['Fix_count_short_66_150', 'Fix_count_medium_300_500',\n",
|
||||
" 'Fix_count_long_gt_1000', 'Fix_count_100', 'Fix_mean_duration',\n",
|
||||
" 'Fix_median_duration', 'Sac_count', 'Sac_mean_amp', 'Sac_mean_dur',\n",
|
||||
" 'Sac_median_dur', 'Pupil_mean', 'Pupil_IPA']\n",
|
||||
"print(len(eye_cols))\n",
|
||||
"all_signal_columns = eye_cols+face_au_cols\n",
|
||||
"print(len(all_signal_columns))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "d5e9c67a",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"MAD"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "592291ef",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def calculate_mad_params(df, columns):\n",
|
||||
" \"\"\"\n",
|
||||
" Calculate median and MAD parameters for each column.\n",
|
||||
" This should be run ONLY on the training data.\n",
|
||||
" \n",
|
||||
" Returns a dictionary: {col: (median, mad)}\n",
|
||||
" \"\"\"\n",
|
||||
" params = {}\n",
|
||||
" for col in columns:\n",
|
||||
" median = df[col].median()\n",
|
||||
" mad = np.median(np.abs(df[col] - median))\n",
|
||||
" params[col] = (median, mad)\n",
|
||||
" return params\n",
|
||||
"def apply_mad_filter(df, params, threshold=3.5):\n",
|
||||
" \"\"\"\n",
|
||||
" Apply MAD-based outlier removal using precomputed parameters.\n",
|
||||
" Works on training, validation, and test data.\n",
|
||||
" \n",
|
||||
" df: DataFrame to filter\n",
|
||||
" params: dictionary {col: (median, mad)} from training data\n",
|
||||
" threshold: cutoff for robust Z-score\n",
|
||||
" \"\"\"\n",
|
||||
" df_clean = df.copy()\n",
|
||||
"\n",
|
||||
" for col, (median, mad) in params.items():\n",
|
||||
" if mad == 0:\n",
|
||||
" continue # no spread; nothing to remove for this column\n",
|
||||
"\n",
|
||||
" robust_z = 0.6745 * (df_clean[col] - median) / mad\n",
|
||||
" outlier_mask = np.abs(robust_z) > threshold\n",
|
||||
"\n",
|
||||
" # Remove values only in this specific column\n",
|
||||
" df_clean.loc[outlier_mask, col] = median\n",
|
||||
" \n",
|
||||
" \n",
|
||||
" print(df_clean.shape)\n",
|
||||
" return df_clean"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4ddad4a8",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"if(FILTER_MAD):\n",
|
||||
" mad_params = calculate_mad_params(df, all_signal_columns)\n",
|
||||
" df = apply_mad_filter(df, mad_params, THRESHOLD)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "89387879",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Normalizer"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9c129cdd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def fit_normalizer(train_data, au_columns, method='standard', scope='global'):\n",
|
||||
" \"\"\"\n",
|
||||
" Fit normalization scalers on training data.\n",
|
||||
" \n",
|
||||
" Parameters:\n",
|
||||
" -----------\n",
|
||||
" train_data : pd.DataFrame\n",
|
||||
" Training dataframe with AU columns and subjectID\n",
|
||||
" au_columns : list\n",
|
||||
" List of AU column names to normalize\n",
|
||||
" method : str, default='standard'\n",
|
||||
" Normalization method: 'standard' for StandardScaler or 'minmax' for MinMaxScaler\n",
|
||||
" scope : str, default='global'\n",
|
||||
" Normalization scope: 'subject' for per-subject or 'global' for across all subjects\n",
|
||||
" \n",
|
||||
" Returns:\n",
|
||||
" --------\n",
|
||||
" dict\n",
|
||||
" Dictionary containing fitted scalers and statistics for new subjects\n",
|
||||
" \"\"\"\n",
|
||||
" if method == 'standard':\n",
|
||||
" Scaler = StandardScaler\n",
|
||||
" elif method == 'minmax':\n",
|
||||
" Scaler = MinMaxScaler\n",
|
||||
" else:\n",
|
||||
" raise ValueError(\"method must be 'standard' or 'minmax'\")\n",
|
||||
" \n",
|
||||
" scalers = {}\n",
|
||||
" if scope == 'subject':\n",
|
||||
" # Fit one scaler per subject\n",
|
||||
" subject_stats = []\n",
|
||||
" \n",
|
||||
" for subject in train_data['subjectID'].unique():\n",
|
||||
" subject_mask = train_data['subjectID'] == subject\n",
|
||||
" scaler = Scaler()\n",
|
||||
" scaler.fit(train_data.loc[subject_mask, au_columns].values)\n",
|
||||
" scalers[subject] = scaler\n",
|
||||
" \n",
|
||||
" # Store statistics for averaging\n",
|
||||
" if method == 'standard':\n",
|
||||
" subject_stats.append({\n",
|
||||
" 'mean': scaler.mean_,\n",
|
||||
" 'std': scaler.scale_\n",
|
||||
" })\n",
|
||||
" elif method == 'minmax':\n",
|
||||
" subject_stats.append({\n",
|
||||
" 'min': scaler.data_min_,\n",
|
||||
" 'max': scaler.data_max_\n",
|
||||
" })\n",
|
||||
" \n",
|
||||
" # Calculate average statistics for new subjects\n",
|
||||
" if method == 'standard':\n",
|
||||
" avg_mean = np.mean([s['mean'] for s in subject_stats], axis=0)\n",
|
||||
" avg_std = np.mean([s['std'] for s in subject_stats], axis=0)\n",
|
||||
" fallback_scaler = StandardScaler()\n",
|
||||
" fallback_scaler.mean_ = avg_mean\n",
|
||||
" fallback_scaler.scale_ = avg_std\n",
|
||||
" fallback_scaler.var_ = avg_std ** 2\n",
|
||||
" fallback_scaler.n_features_in_ = len(au_columns)\n",
|
||||
" elif method == 'minmax':\n",
|
||||
" avg_min = np.mean([s['min'] for s in subject_stats], axis=0)\n",
|
||||
" avg_max = np.mean([s['max'] for s in subject_stats], axis=0)\n",
|
||||
" fallback_scaler = MinMaxScaler()\n",
|
||||
" fallback_scaler.data_min_ = avg_min\n",
|
||||
" fallback_scaler.data_max_ = avg_max\n",
|
||||
" fallback_scaler.data_range_ = avg_max - avg_min\n",
|
||||
" fallback_scaler.scale_ = 1.0 / fallback_scaler.data_range_\n",
|
||||
" fallback_scaler.min_ = -avg_min * fallback_scaler.scale_\n",
|
||||
" fallback_scaler.n_features_in_ = len(au_columns)\n",
|
||||
" \n",
|
||||
" scalers['_fallback'] = fallback_scaler\n",
|
||||
" \n",
|
||||
" elif scope == 'global':\n",
|
||||
" # Fit one scaler for all subjects\n",
|
||||
" scaler = Scaler()\n",
|
||||
" scaler.fit(train_data[au_columns].values)\n",
|
||||
" scalers['global'] = scaler\n",
|
||||
" \n",
|
||||
" else:\n",
|
||||
" raise ValueError(\"scope must be 'subject' or 'global'\")\n",
|
||||
" \n",
|
||||
" return {'scalers': scalers, 'method': method, 'scope': scope}"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9cfabd37",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def apply_normalizer(data, columns, normalizer_dict):\n",
|
||||
" \"\"\"\n",
|
||||
" Apply fitted normalization scalers to data.\n",
|
||||
" \n",
|
||||
" Parameters:\n",
|
||||
" -----------\n",
|
||||
" data : pd.DataFrame\n",
|
||||
" Dataframe with AU columns and subjectID\n",
|
||||
" au_columns : list\n",
|
||||
" List of AU column names to normalize\n",
|
||||
" normalizer_dict : dict\n",
|
||||
" Dictionary containing fitted scalers from fit_normalizer()\n",
|
||||
" \n",
|
||||
" Returns:\n",
|
||||
" --------\n",
|
||||
" pd.DataFrame\n",
|
||||
" DataFrame with normalized AU columns\n",
|
||||
" \"\"\"\n",
|
||||
" normalized_data = data.copy()\n",
|
||||
" scalers = normalizer_dict['scalers']\n",
|
||||
" scope = normalizer_dict['scope']\n",
|
||||
" normalized_data[columns] = normalized_data[columns].astype(np.float64)\n",
|
||||
"\n",
|
||||
" if scope == 'subject':\n",
|
||||
" # Apply per-subject normalization\n",
|
||||
" for subject in data['subjectID'].unique():\n",
|
||||
" subject_mask = data['subjectID'] == subject\n",
|
||||
" \n",
|
||||
" # Use the subject's scaler if available, otherwise use fallback\n",
|
||||
" if subject in scalers:\n",
|
||||
" scaler = scalers[subject]\n",
|
||||
" else:\n",
|
||||
" # Use averaged scaler for new subjects\n",
|
||||
" scaler = scalers['_fallback']\n",
|
||||
" print(f\"Info: Subject {subject} not in training data. Using averaged scaler from training subjects.\")\n",
|
||||
" \n",
|
||||
" normalized_data.loc[subject_mask, columns] = scaler.transform(\n",
|
||||
" data.loc[subject_mask, columns].values\n",
|
||||
" )\n",
|
||||
" \n",
|
||||
" elif scope == 'global':\n",
|
||||
" # Apply global normalization\n",
|
||||
" scaler = scalers['global']\n",
|
||||
" normalized_data[columns] = scaler.transform(data[columns].values)\n",
|
||||
" \n",
|
||||
" return normalized_data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4dbbebf7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"scaler = fit_normalizer(df, all_signal_columns, method=METHOD, scope=SCOPE)\n",
|
||||
"df_min_max_normalised = apply_normalizer(df, all_signal_columns, scaler)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "6b9b2ae8",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"a= df_min_max_normalised[['STUDY','LEVEL','PHASE']]\n",
|
||||
"print(a.dtypes)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e3e1bc34",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Define signal columns (adjust only once)\n",
|
||||
"signal_columns = all_signal_columns\n",
|
||||
"\n",
|
||||
"# Get all unique combinations of STUDY, LEVEL and PHASE\n",
|
||||
"unique_combinations = df_min_max_normalised[['STUDY', 'LEVEL', 'PHASE']].drop_duplicates().reset_index(drop=True)\n",
|
||||
"\n",
|
||||
"# Dictionary to store subsets\n",
|
||||
"subsets = {}\n",
|
||||
"subset_sizes = {}\n",
|
||||
"\n",
|
||||
"for idx, row in unique_combinations.iterrows():\n",
|
||||
" study = row['STUDY']\n",
|
||||
" level = row['LEVEL']\n",
|
||||
" phase = row['PHASE']\n",
|
||||
" key = f\"{study}_L{level}_P{phase}\"\n",
|
||||
" subset = df_min_max_normalised[\n",
|
||||
" (df_min_max_normalised['STUDY'] == study) & \n",
|
||||
" (df_min_max_normalised['LEVEL'] == level) & \n",
|
||||
" (df_min_max_normalised['PHASE'] == phase)\n",
|
||||
" ]\n",
|
||||
" subsets[key] = subset\n",
|
||||
" subset_sizes[key] = len(subset)\n",
|
||||
"\n",
|
||||
"# Output subset sizes\n",
|
||||
"print(\"Number of samples per subset:\")\n",
|
||||
"print(\"=\" * 40)\n",
|
||||
"for key, size in subset_sizes.items():\n",
|
||||
" print(f\"{key}: {size} samples\")\n",
|
||||
"print(\"=\" * 40)\n",
|
||||
"print(f\"Total number of subsets: {len(subsets)}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "c7fdeb5c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"\n",
|
||||
"# Function to categorize subsets\n",
|
||||
"def categorize_subset(key):\n",
|
||||
" \"\"\"Categorizes a subset as 'low' or 'high' based on the given logic\"\"\"\n",
|
||||
" parts = key.split('_')\n",
|
||||
" study = parts[0]\n",
|
||||
" level = int(parts[1][1:]) # 'L1' -> 1\n",
|
||||
" phase = parts[2][1:] # 'Pbaseline' -> 'baseline'\n",
|
||||
" \n",
|
||||
" # LOW: baseline OR (n-back with level 1 or 4)\n",
|
||||
" if phase == \"baseline\":\n",
|
||||
" return 'low'\n",
|
||||
" elif study == \"n-back\" and level in [1, 4]:\n",
|
||||
" return 'low'\n",
|
||||
" \n",
|
||||
" # HIGH: (n-back with level 2,3,5,6 and phase train/test) OR (k-drive not baseline)\n",
|
||||
" elif study == \"n-back\" and level in [2, 3, 5, 6] and phase in [\"train\", \"test\"]:\n",
|
||||
" return 'high'\n",
|
||||
" elif study == \"k-drive\" and phase != \"baseline\":\n",
|
||||
" return 'high'\n",
|
||||
" \n",
|
||||
" return None\n",
|
||||
"\n",
|
||||
"# Categorize subsets\n",
|
||||
"low_subsets = {}\n",
|
||||
"high_subsets = {}\n",
|
||||
"\n",
|
||||
"for key, subset in subsets.items():\n",
|
||||
" category = categorize_subset(key)\n",
|
||||
" if category == 'low':\n",
|
||||
" low_subsets[key] = subset\n",
|
||||
" elif category == 'high':\n",
|
||||
" high_subsets[key] = subset\n",
|
||||
"\n",
|
||||
"# Output statistics\n",
|
||||
"print(\"\\n\" + \"=\" * 50)\n",
|
||||
"print(\"SUBSET CATEGORIZATION\")\n",
|
||||
"print(\"=\" * 50)\n",
|
||||
"\n",
|
||||
"print(\"\\nLOW subsets (Blue):\")\n",
|
||||
"print(\"-\" * 50)\n",
|
||||
"low_total = 0\n",
|
||||
"for key in sorted(low_subsets.keys()):\n",
|
||||
" size = subset_sizes[key]\n",
|
||||
" low_total += size\n",
|
||||
" print(f\" {key}: {size} samples\")\n",
|
||||
"print(f\"{'TOTAL LOW:':<30} {low_total} samples\")\n",
|
||||
"print(f\"{'NUMBER OF LOW SUBSETS:':<30} {len(low_subsets)}\")\n",
|
||||
"\n",
|
||||
"print(\"\\nHIGH subsets (Red):\")\n",
|
||||
"print(\"-\" * 50)\n",
|
||||
"high_total = 0\n",
|
||||
"for key in sorted(high_subsets.keys()):\n",
|
||||
" size = subset_sizes[key]\n",
|
||||
" high_total += size\n",
|
||||
" print(f\" {key}: {size} samples\")\n",
|
||||
"print(f\"{'TOTAL HIGH:':<30} {high_total} samples\")\n",
|
||||
"print(f\"{'NUMBER OF HIGH SUBSETS:':<30} {len(high_subsets)}\")\n",
|
||||
"\n",
|
||||
"print(\"\\n\" + \"=\" * 50)\n",
|
||||
"print(f\"TOTAL SAMPLES: {low_total + high_total}\")\n",
|
||||
"print(f\"TOTAL SUBSETS: {len(low_subsets) + len(high_subsets)}\")\n",
|
||||
"print(\"=\" * 50)\n",
|
||||
"\n",
|
||||
"# Find minimum subset size\n",
|
||||
"min_subset_size = min(subset_sizes.values())\n",
|
||||
"print(f\"\\nMinimum subset size: {min_subset_size}\")\n",
|
||||
"\n",
|
||||
"# Number of points to plot per subset (50% of minimum size)\n",
|
||||
"sampling_factor = 1\n",
|
||||
"n_samples_per_subset = int(sampling_factor * min_subset_size)\n",
|
||||
"print(f\"Number of randomly drawn points per subset: {n_samples_per_subset}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "ff363fc5",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Plot"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "3a9d9163",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Create comparison plots\n",
|
||||
"fig, axes = plt.subplots(len(signal_columns), 1, figsize=(14, 4 * len(signal_columns)))\n",
|
||||
"\n",
|
||||
"# If only one signal column exists, convert axes to list\n",
|
||||
"if len(signal_columns) == 1:\n",
|
||||
" axes = [axes]\n",
|
||||
"\n",
|
||||
"# Create a plot for each signal column\n",
|
||||
"for i, signal_col in enumerate(signal_columns):\n",
|
||||
" ax = axes[i]\n",
|
||||
" \n",
|
||||
" y_pos = 0\n",
|
||||
" labels = []\n",
|
||||
" \n",
|
||||
" # First plot all LOW subsets (sorted, blue)\n",
|
||||
" for label in sorted(low_subsets.keys()):\n",
|
||||
" subset = low_subsets[label]\n",
|
||||
" if len(subset) > 0 and signal_col in subset.columns:\n",
|
||||
" # Draw random sample\n",
|
||||
" n_samples = min(n_samples_per_subset, len(subset))\n",
|
||||
" sampled_data = subset[signal_col].sample(n=n_samples, random_state=42)\n",
|
||||
" \n",
|
||||
" # Calculate mean and median\n",
|
||||
" mean_val = subset[signal_col].mean()\n",
|
||||
" median_val = subset[signal_col].median()\n",
|
||||
" \n",
|
||||
" # Plot points in blue\n",
|
||||
" ax.scatter(sampled_data, [y_pos] * len(sampled_data), \n",
|
||||
" alpha=0.5, s=30, color='blue')\n",
|
||||
" \n",
|
||||
" # Mean as black cross\n",
|
||||
" ax.plot(mean_val, y_pos, 'x', markersize=12, markeredgewidth=3, \n",
|
||||
" color='black', zorder=5)\n",
|
||||
" \n",
|
||||
" # Median as brown cross\n",
|
||||
" ax.plot(median_val, y_pos, 'x', markersize=12, markeredgewidth=3, \n",
|
||||
" color='brown', zorder=5)\n",
|
||||
" \n",
|
||||
" labels.append(f\"{label} (n={subset_sizes[label]})\")\n",
|
||||
" y_pos += 1\n",
|
||||
" \n",
|
||||
" # Separation line between LOW and HIGH\n",
|
||||
" if len(low_subsets) > 0 and len(high_subsets) > 0:\n",
|
||||
" ax.axhline(y=y_pos - 0.5, color='gray', linestyle='--', linewidth=2, alpha=0.7)\n",
|
||||
" \n",
|
||||
" # Then plot all HIGH subsets (sorted, red)\n",
|
||||
" for label in sorted(high_subsets.keys()):\n",
|
||||
" subset = high_subsets[label]\n",
|
||||
" if len(subset) > 0 and signal_col in subset.columns:\n",
|
||||
" # Draw random sample\n",
|
||||
" n_samples = min(n_samples_per_subset, len(subset))\n",
|
||||
" sampled_data = subset[signal_col].sample(n=n_samples, random_state=42)\n",
|
||||
" \n",
|
||||
" # Calculate mean and median\n",
|
||||
" mean_val = subset[signal_col].mean()\n",
|
||||
" median_val = subset[signal_col].median()\n",
|
||||
" \n",
|
||||
" # Plot points in red\n",
|
||||
" ax.scatter(sampled_data, [y_pos] * len(sampled_data), \n",
|
||||
" alpha=0.5, s=30, color='red')\n",
|
||||
" \n",
|
||||
" # Mean as black cross\n",
|
||||
" ax.plot(mean_val, y_pos, 'x', markersize=12, markeredgewidth=3, \n",
|
||||
" color='black', zorder=5)\n",
|
||||
" \n",
|
||||
" # Median as brown cross\n",
|
||||
" ax.plot(median_val, y_pos, 'x', markersize=12, markeredgewidth=3, \n",
|
||||
" color='brown', zorder=5)\n",
|
||||
" \n",
|
||||
" labels.append(f\"{label} (n={subset_sizes[label]})\")\n",
|
||||
" y_pos += 1\n",
|
||||
" \n",
|
||||
" ax.set_yticks(range(len(labels)))\n",
|
||||
" ax.set_yticklabels(labels)\n",
|
||||
" ax.set_xlabel(f'{signal_col} value')\n",
|
||||
" ax.set_title(f'{signal_col}: LOW (Blue) vs HIGH (Red) | {n_samples_per_subset} points/subset | Black X = Mean, Brown X = Median')\n",
|
||||
" ax.grid(True, alpha=0.3, axis='x')\n",
|
||||
" ax.axvline(0, color='gray', linestyle='--', alpha=0.5)\n",
|
||||
"\n",
|
||||
"plt.tight_layout()\n",
|
||||
"plt.show()\n",
|
||||
"\n",
|
||||
"print(f\"\\nNote: {n_samples_per_subset} random points were plotted per subset.\")\n",
|
||||
"print(\"Blue points = LOW subsets | Red points = HIGH subsets\")\n",
|
||||
"print(\"Black 'X' = Mean of entire subset | Brown 'X' = Median of entire subset\")\n",
|
||||
"print(f\"Total subsets plotted: {len(low_subsets)} LOW + {len(high_subsets)} HIGH = {len(low_subsets) + len(high_subsets)} subsets\")"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
+55
-22
@@ -1,5 +1,13 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "cc08936c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Insights into the dataset with histogramms and scatter plots"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "1014c5e0",
|
||||
@@ -17,7 +25,8 @@
|
||||
"source": [
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"import matplotlib.pyplot as plt"
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"from pathlib import Path"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -27,7 +36,7 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"path =r\"C:\\Users\\micha\\FAUbox\\WS2526_Fahrsimulator_MSY (Celina Korzer)\\AU_dataset\\output_windowed.parquet\"\n",
|
||||
"path = Path(r\".parquet\") # TODO: enter path to dataset\n",
|
||||
"df = pd.read_parquet(path=path)"
|
||||
]
|
||||
},
|
||||
@@ -104,21 +113,27 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Get all columns that start with 'AU'\n",
|
||||
"au_columns = [col for col in low_all.columns if col.startswith('AU')]\n",
|
||||
"face_au_cols = [c for c in low_all.columns if c.startswith(\"FACE_AU\")]\n",
|
||||
"eye_cols = ['Fix_count_short_66_150', 'Fix_count_medium_300_500',\n",
|
||||
" 'Fix_count_long_gt_1000', 'Fix_count_100', 'Fix_mean_duration',\n",
|
||||
" 'Fix_median_duration', 'Sac_count', 'Sac_mean_amp', 'Sac_mean_dur',\n",
|
||||
" 'Sac_median_dur', 'Blink_count', 'Blink_mean_dur', 'Blink_median_dur',\n",
|
||||
" 'Pupil_mean', 'Pupil_IPA']\n",
|
||||
"\n",
|
||||
"cols = face_au_cols+eye_cols\n",
|
||||
"\n",
|
||||
"# Calculate number of rows and columns for subplots\n",
|
||||
"n_cols = len(au_columns)\n",
|
||||
"n_rows = 4\n",
|
||||
"n_cols = len(cols)\n",
|
||||
"n_rows = 7\n",
|
||||
"n_cols_subplot = 5\n",
|
||||
"\n",
|
||||
"# Create figure with subplots\n",
|
||||
"fig, axes = plt.subplots(n_rows, n_cols_subplot, figsize=(20, 16))\n",
|
||||
"axes = axes.flatten()\n",
|
||||
"fig.suptitle('Action Unit (AU) Distributions: Low vs High', fontsize=20, fontweight='bold', y=0.995)\n",
|
||||
"fig.suptitle('Feature Distributions: Low vs High', fontsize=20, fontweight='bold', y=0.995)\n",
|
||||
"\n",
|
||||
"# Create histogram for each AU column\n",
|
||||
"for idx, col in enumerate(au_columns):\n",
|
||||
"for idx, col in enumerate(cols):\n",
|
||||
" ax = axes[idx]\n",
|
||||
" \n",
|
||||
" # Plot overlapping histograms\n",
|
||||
@@ -133,32 +148,50 @@
|
||||
" ax.grid(True, alpha=0.3)\n",
|
||||
"\n",
|
||||
"# Hide any unused subplots\n",
|
||||
"for idx in range(len(au_columns), len(axes)):\n",
|
||||
"for idx in range(len(cols), len(axes)):\n",
|
||||
" axes[idx].set_visible(False)\n",
|
||||
"\n",
|
||||
"# Adjust layout\n",
|
||||
"plt.tight_layout()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "6cd53cdb",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Create figure with subplots\n",
|
||||
"fig, axes = plt.subplots(n_rows, n_cols_subplot, figsize=(20, 16))\n",
|
||||
"axes = axes.flatten()\n",
|
||||
"fig.suptitle('Feature Scatter: Low vs High', fontsize=20, fontweight='bold', y=0.995)\n",
|
||||
"\n",
|
||||
"for idx, col in enumerate(cols):\n",
|
||||
" ax = axes[idx]\n",
|
||||
"\n",
|
||||
" # Scatterplots\n",
|
||||
" ax.scatter(range(len(low_all[col])), low_all[col], alpha=0.6, color='blue', label='low_all', s=10)\n",
|
||||
" ax.scatter(range(len(high_all[col])), high_all[col], alpha=0.6, color='red', label='high_all', s=10)\n",
|
||||
"\n",
|
||||
" ax.set_title(col, fontsize=10, fontweight='bold')\n",
|
||||
" ax.set_xlabel('Sample index', fontsize=8)\n",
|
||||
" ax.set_ylabel('Value', fontsize=8)\n",
|
||||
" ax.legend(fontsize=8)\n",
|
||||
" ax.grid(True, alpha=0.3)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"plt.tight_layout()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "base",
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.11.5"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
- url: # enter url
|
||||
- password: # enter passwort
|
||||
@@ -1,157 +0,0 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "aab6b326-a583-47ad-8bb7-723c2fddcc63",
|
||||
"metadata": {
|
||||
"scrolled": true
|
||||
},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# %pip install pyocclient\n",
|
||||
"import yaml\n",
|
||||
"import owncloud\n",
|
||||
"import pandas as pd\n",
|
||||
"import time"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4f42846c-27c3-4394-a40a-e22d73c2902e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"start = time.time()\n",
|
||||
"\n",
|
||||
"with open(\"../login.yaml\") as f:\n",
|
||||
" cfg = yaml.safe_load(f)\n",
|
||||
"url, password = cfg[0][\"url\"], cfg[1][\"password\"]\n",
|
||||
"file = \"adabase-public-0022-v_0_0_2.h5py\"\n",
|
||||
"oc = owncloud.Client.from_public_link(url, folder_password=password)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"oc.get_file(file, \"tmp22.h5\")\n",
|
||||
"\n",
|
||||
"end = time.time()\n",
|
||||
"print(end - start)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "3714dec2-85d0-4f76-af46-ea45ebec2fa3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"start = time.time()\n",
|
||||
"df_performance = pd.read_hdf(\"tmp22.h5\", \"PERFORMANCE\")\n",
|
||||
"end = time.time()\n",
|
||||
"print(end - start)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "f50e97d0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"print(22)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "c131c816",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_performance"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "6ae47e52-ad86-4f8d-b929-0080dc99f646",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"start = time.time()\n",
|
||||
"df_4_col = pd.read_hdf(\"tmp.h5\", \"SIGNALS\", mode=\"r\", columns=[\"STUDY\"], start=0, stop=1)\n",
|
||||
"end = time.time()\n",
|
||||
"print(end - start)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "7c139f3a-ede8-4530-957d-d1bb939f6cb5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_4_col.head()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "a68d58ea-65f2-46c4-a2b2-8c3447c715d7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_4_col.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "95aa4523-3784-4ab6-bf92-0227ce60e863",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_4_col.info()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "defbcaf4-ad1b-453f-9b48-ab0ecfc4b5d5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_4_col.isna().sum()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "72313895-c478-44a5-9108-00b0bec01bb8",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": []
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "base",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.11.5"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,112 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "457e7807",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"# Get data from owncloud"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "dc9ed3f8",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Imports"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# %pip install pyocclient\n",
|
||||
"import os\n",
|
||||
"import time\n",
|
||||
"import yaml\n",
|
||||
"import owncloud\n",
|
||||
"import pandas as pd\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "68e34abc",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Download and save"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"start = time.time()\n",
|
||||
"\n",
|
||||
"# TODO: User input: directory where downloaded files should be saved\n",
|
||||
"save_dir = r\"./downloads\"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"os.makedirs(save_dir, exist_ok=True)\n",
|
||||
"\n",
|
||||
"# Load credentials\n",
|
||||
"with open(\"login.yaml\", \"r\") as f:\n",
|
||||
" cfg = yaml.safe_load(f)\n",
|
||||
"\n",
|
||||
"url = cfg[0][\"url\"]\n",
|
||||
"password = cfg[1][\"password\"]\n",
|
||||
"\n",
|
||||
"# Connect to OwnCloud public link\n",
|
||||
"oc = owncloud.Client.from_public_link(url, folder_password=password)\n",
|
||||
"\n",
|
||||
"# List all available files in the shared folder\n",
|
||||
"remote_files = oc.list(\".\")\n",
|
||||
"\n",
|
||||
"# Keep only HDF5 files and sort them by name\n",
|
||||
"hdf5_files = sorted([f.get_name() for f in remote_files if f.get_name().endswith(\".hdf5\")])\n",
|
||||
"\n",
|
||||
"print(f\"Found {len(hdf5_files)} .hdf5 files in OwnCloud\")\n",
|
||||
"\n",
|
||||
"if not hdf5_files:\n",
|
||||
" print(\"No .hdf5 files found.\")\n",
|
||||
"else:\n",
|
||||
" for i, remote_name in enumerate(hdf5_files):\n",
|
||||
" local_name = f\"tmp_{i:04d}.h5\"\n",
|
||||
" local_path = os.path.join(save_dir, local_name)\n",
|
||||
"\n",
|
||||
" try:\n",
|
||||
" oc.get_file(remote_name, local_path)\n",
|
||||
" print(f\"Downloaded: {remote_name} -> {local_path}\")\n",
|
||||
" except Exception as e:\n",
|
||||
" print(f\"Failed to download {remote_name}: {e}\")\n",
|
||||
"\n",
|
||||
"end = time.time()\n",
|
||||
"print(f\"Finished in {end - start:.2f} seconds\")\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "base",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.11.5"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -15,6 +15,7 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"%pip install pyocclient\n",
|
||||
"import yaml\n",
|
||||
"import owncloud\n",
|
||||
"import pandas as pd\n",
|
||||
@@ -36,101 +37,109 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Load credentials\n",
|
||||
"with open(\"../login.yaml\") as f:\n",
|
||||
"# Load credentials from YAML\n",
|
||||
"with open(\"login.yaml\", \"r\") as f:\n",
|
||||
" cfg = yaml.safe_load(f)\n",
|
||||
" \n",
|
||||
"url, password = cfg[0][\"url\"], cfg[1][\"password\"]\n",
|
||||
"\n",
|
||||
"# Connect once\n",
|
||||
"url = cfg[0][\"url\"]\n",
|
||||
"password = cfg[1][\"password\"]\n",
|
||||
"\n",
|
||||
"# Connect once to the public OwnCloud link\n",
|
||||
"oc = owncloud.Client.from_public_link(url, folder_password=password)\n",
|
||||
"# File pattern\n",
|
||||
"# base = \"adabase-public-{num:04d}-v_0_0_2.h5py\"\n",
|
||||
"base = \"{num:04d}-*.h5py\""
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "07c03d07",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"num_files = 2 # number of files to process (min: 1, max: 30)\n",
|
||||
"\n",
|
||||
"num_files = 1 # number of subject IDs to process (min: 1, max: 30)\n",
|
||||
"performance_data = []\n",
|
||||
"\n",
|
||||
"# Read remote file list once\n",
|
||||
"remote_files = oc.list(\".\")\n",
|
||||
"remote_names = [f.get_name() for f in remote_files]\n",
|
||||
"\n",
|
||||
"for i in range(num_files):\n",
|
||||
" file_pattern = f\"{i:04d}-*\"\n",
|
||||
" \n",
|
||||
" # Get list of files matching the pattern\n",
|
||||
" files = oc.list('.')\n",
|
||||
" matching_files = [f.get_name() for f in files if f.get_name().startswith(f\"{i:04d}-\")]\n",
|
||||
" \n",
|
||||
" if matching_files:\n",
|
||||
" file_name = matching_files[0] # Take the first matching file\n",
|
||||
" local_tmp = f\"tmp_{i:04d}.h5\"\n",
|
||||
" \n",
|
||||
" oc.get_file(file_name, local_tmp)\n",
|
||||
" print(f\"{file_name} geöffnet\")\n",
|
||||
" else:\n",
|
||||
" print(f\"Keine Datei gefunden für Muster: {file_pattern}\")\n",
|
||||
" # file_name = base.format(num=i)\n",
|
||||
" # local_tmp = f\"tmp_{i:04d}.h5\"\n",
|
||||
" prefix = f\"{i:04d}-\"\n",
|
||||
" matching_files = [name for name in remote_names if name.startswith(prefix) and name.endswith(\".hdf5\")]\n",
|
||||
"\n",
|
||||
" # oc.get_file(file_name, local_tmp)\n",
|
||||
" # print(f\"{file_name} geöffnet\")\n",
|
||||
"\n",
|
||||
" # check SIGNALS table for AUs\n",
|
||||
" with pd.HDFStore(local_tmp, mode=\"r\") as store:\n",
|
||||
" cols = store.select(\"SIGNALS\", start=0, stop=1).columns\n",
|
||||
" au_cols = [c for c in cols if c.startswith(\"AU\")]\n",
|
||||
" if not au_cols:\n",
|
||||
" print(f\"Subject {i} enthält keine AUs\")\n",
|
||||
" if not matching_files:\n",
|
||||
" print(f\"No file found for pattern: {prefix}*.hdf5\")\n",
|
||||
" continue\n",
|
||||
"\n",
|
||||
" # load performance table\n",
|
||||
" # Take the first matching file, e.g. 0000-AACA.hdf5\n",
|
||||
" file_name = matching_files[0]\n",
|
||||
" local_tmp = f\"tmp_{i:04d}.hdf5\"\n",
|
||||
"\n",
|
||||
" try:\n",
|
||||
" # Download the file locally\n",
|
||||
" oc.get_file(file_name, local_tmp)\n",
|
||||
" print(f\"Downloaded and opened file: {file_name} -> {local_tmp}\")\n",
|
||||
" except Exception as e:\n",
|
||||
" print(f\"Failed to download file {file_name}: {e}\")\n",
|
||||
" continue\n",
|
||||
"\n",
|
||||
" # Check SIGNALS table for AU columns\n",
|
||||
" try:\n",
|
||||
" with pd.HDFStore(local_tmp, mode=\"r\") as store:\n",
|
||||
" cols = store.select(\"SIGNALS\", start=0, stop=1).columns\n",
|
||||
" except Exception as e:\n",
|
||||
" print(f\"Failed to read SIGNALS from {local_tmp}: {e}\")\n",
|
||||
" continue\n",
|
||||
"\n",
|
||||
" au_cols = [c for c in cols if c.startswith(\"AU\")]\n",
|
||||
" if not au_cols:\n",
|
||||
" print(f\"Subject {i:04d} contains no AU columns\")\n",
|
||||
" continue\n",
|
||||
"\n",
|
||||
" # Load PERFORMANCE table\n",
|
||||
" try:\n",
|
||||
" with pd.HDFStore(local_tmp, mode=\"r\") as store:\n",
|
||||
" perf_df = store.select(\"PERFORMANCE\")\n",
|
||||
" except Exception as e:\n",
|
||||
" print(f\"Failed to read PERFORMANCE from {local_tmp}: {e}\")\n",
|
||||
" continue\n",
|
||||
"\n",
|
||||
" f1_cols = [c for c in [\"AUDITIVE F1\", \"VISUAL F1\", \"F1\"] if c in perf_df.columns]\n",
|
||||
" if not f1_cols:\n",
|
||||
" print(f\"Subject {i}: keine F1-Spalten gefunden\")\n",
|
||||
" print(f\"Subject {i:04d}: no F1 columns found\")\n",
|
||||
" continue\n",
|
||||
"\n",
|
||||
" subject_entry = {\"subjectID\": i}\n",
|
||||
" valid_scores = []\n",
|
||||
"\n",
|
||||
" # iterate rows: each (study, level, phase)\n",
|
||||
" # Iterate through PERFORMANCE rows: each row is one (study, level, phase) combination\n",
|
||||
" for _, row in perf_df.iterrows():\n",
|
||||
" study, level, phase = row[\"STUDY\"], row[\"LEVEL\"], row[\"PHASE\"]\n",
|
||||
" study = row[\"STUDY\"]\n",
|
||||
" level = row[\"LEVEL\"]\n",
|
||||
" phase = row[\"PHASE\"]\n",
|
||||
" col_name = f\"STUDY_{study}_LEVEL_{level}_PHASE_{phase}\"\n",
|
||||
"\n",
|
||||
" # collect valid F1 values among the three columns\n",
|
||||
" # Collect non-NaN F1 values from the available F1 columns\n",
|
||||
" scores = [row[c] for c in f1_cols if pd.notna(row[c])]\n",
|
||||
" if scores:\n",
|
||||
" mean_score = float(np.mean(scores))\n",
|
||||
" subject_entry[col_name] = mean_score\n",
|
||||
" valid_scores.extend(scores)\n",
|
||||
"\n",
|
||||
" # compute overall average across all valid combinations\n",
|
||||
" # Compute overall average across all valid F1 values\n",
|
||||
" if valid_scores:\n",
|
||||
" subject_entry[\"overall_score\"] = float(np.mean(valid_scores))\n",
|
||||
" performance_data.append(subject_entry)\n",
|
||||
" print(f\"Subject {i}: {len(valid_scores)} gültige Scores, Overall = {subject_entry['overall_score']:.3f}\")\n",
|
||||
" print(\n",
|
||||
" f\"Subject {i:04d}: {len(valid_scores)} valid scores, \"\n",
|
||||
" f\"overall = {subject_entry['overall_score']:.3f}\"\n",
|
||||
" )\n",
|
||||
" else:\n",
|
||||
" print(f\"Subject {i}: keine gültigen F1-Scores\")\n",
|
||||
" print(f\"Subject {i:04d}: no valid F1 scores found\")\n",
|
||||
"\n",
|
||||
"# build dataframe\n",
|
||||
"# Build final DataFrame and save CSV\n",
|
||||
"if performance_data:\n",
|
||||
" performance_df = pd.DataFrame(performance_data)\n",
|
||||
" combination_cols = sorted([c for c in performance_df.columns if c.startswith(\"STUDY_\")])\n",
|
||||
" final_cols = [\"subjectID\", \"overall_score\"] + combination_cols\n",
|
||||
" performance_df = performance_df[final_cols]\n",
|
||||
" performance_df.to_csv(\"n_au_performance.csv\", index=False)\n",
|
||||
" performance_df.to_csv(\"performance.csv\", index=False)\n",
|
||||
"\n",
|
||||
" print(f\"\\nGesamt Subjects mit Action Units: {len(performance_df)}\")\n",
|
||||
" print(f\"\\nTotal subjects with Action Units: {len(performance_df)}\")\n",
|
||||
" print(\"Saved results to performance.csv\")\n",
|
||||
"else:\n",
|
||||
" print(\"Keine gültigen Daten gefunden.\")"
|
||||
" print(\"No valid data found.\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -142,56 +151,11 @@
|
||||
"source": [
|
||||
"performance_df.head()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "db95eea7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"with pd.HDFStore(\"tmp_0000.h5\", mode=\"r\") as store:\n",
|
||||
" md = store.select(\"META\")\n",
|
||||
"print(\"File 0:\")\n",
|
||||
"print(md)\n",
|
||||
"with pd.HDFStore(\"tmp_0001.h5\", mode=\"r\") as store:\n",
|
||||
" md = store.select(\"META\")\n",
|
||||
"print(\"File 1\")\n",
|
||||
"print(md)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "8067036b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"pd.set_option('display.max_columns', None)\n",
|
||||
"pd.set_option('display.max_rows', None)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "f18e7385",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"with pd.HDFStore(\"tmp_0000.h5\", mode=\"r\") as store:\n",
|
||||
" md = store.select(\"SIGNALS\", start=0, stop=1)\n",
|
||||
"print(\"File 0:\")\n",
|
||||
"md.head()\n",
|
||||
"# with pd.HDFStore(\"tmp_0001.h5\", mode=\"r\",start=0, stop=1) as store:\n",
|
||||
"# md = store.select(\"SIGNALS\")\n",
|
||||
"# print(\"File 1\")\n",
|
||||
"# print(md.columns)"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "base",
|
||||
"display_name": "310",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
@@ -205,7 +169,7 @@
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.11.5"
|
||||
"version": "3.10.19"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
|
||||
@@ -1,58 +0,0 @@
|
||||
from feat import Detector
|
||||
from feat.utils.io import get_test_data_path
|
||||
from moviepy.video.io.VideoFileClip import VideoFileClip
|
||||
import os
|
||||
|
||||
def extract_aus(path, model):
|
||||
detector = Detector(au_model=model)
|
||||
|
||||
video_prediction = detector.detect(
|
||||
path, data_type="video", skip_frames=24*5, face_detection_threshold=0.95 # alle 5 Sekunden einbeziehen - 24 Frames pro Sekunde
|
||||
)
|
||||
|
||||
return video_prediction.aus.sum()
|
||||
|
||||
def split_video(path, chunk_length=120):
|
||||
video = VideoFileClip(path)
|
||||
duration = int(video.duration)
|
||||
|
||||
subclips_dir = os.path.join(os.dirname(path), "subclips")
|
||||
os.makedirs(subclips_dir, exist_ok=True)
|
||||
paths = []
|
||||
|
||||
for start in range(0, duration, chunk_length):
|
||||
end = min(start + chunk_length, duration)
|
||||
|
||||
subclip = (
|
||||
video
|
||||
.subclip(start, end)
|
||||
.without_audio()
|
||||
.set_fps(video.fps)
|
||||
)
|
||||
|
||||
output_path = f"{subclips_dir}_part_{start//chunk_length + 1}.mp4"
|
||||
subclip.write_videofile(
|
||||
output_path,
|
||||
)
|
||||
paths.append(output_path)
|
||||
|
||||
return output_path
|
||||
|
||||
def start(path):
|
||||
results = []
|
||||
clips = split_video(path)
|
||||
|
||||
for clip in clips:
|
||||
results.append(extract_aus(clip, 'svm'))
|
||||
return results
|
||||
|
||||
if __name__ == "__main__":
|
||||
results = []
|
||||
clips = []
|
||||
test_video_path = "AU_creation/YTDown.com_YouTube_Was-ist-los-bei-7-vs-Wild_Media_Gtj9zu_WikU_001_1080p.mp4"
|
||||
clips = split_video(test_video_path)
|
||||
|
||||
for clippath in clips:
|
||||
results.append(extract_aus(clippath, 'svm'))
|
||||
|
||||
print(results)
|
||||
@@ -5,27 +5,47 @@
|
||||
"id": "3b0c6c82",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Hier entsteht die Dokumentation, wie die Action Units erzeugt wurden.\n",
|
||||
"Daraus wird dann letztendlich ein Skript erstellt, welches automatisch AUs aus Videodateien erstellen soll.\n",
|
||||
"## Action Unit Documentation and Setup\n",
|
||||
"\n",
|
||||
"Py-Feat besitzt Dependencies, die ab Python 3.12 nicht mehr verfügbar sind.\n",
|
||||
"Dazu muss ein Kernel mit Python 3.11 erstellt werden.\n",
|
||||
"Folgendes Vorgehen:\n",
|
||||
"1. Seite des Jupyter Labs öffnen\n",
|
||||
"2. Terminal öffnen und folgende Befehle eingeben:\n",
|
||||
" conda create -n py311 python=3.11\n",
|
||||
" source ~/.bashrc\n",
|
||||
" conda activate py311\n",
|
||||
" conda install jupyter\n",
|
||||
" python -m ipykernel install --user --name=py311 --display-name \"Python 3.11\"\n",
|
||||
" pip install py-feat\n",
|
||||
" pip install \"moviepy<2.0\" (falls benötigt)\n",
|
||||
"3. den Kernel neustarten\n",
|
||||
"4. in VSC den Kernel neu hinzufügen und dann den Kernel mit dem Namen \"Python 3.11\" auswählen.\n",
|
||||
"This documentation outlines the process for generating **Action Units (AUs)** and the eventual creation of a script to automate AU extraction from video files.\n",
|
||||
"\n",
|
||||
"Der Code unten zeigt eine beispielhafte Integration der py-feat Bibliothek.\n",
|
||||
"Die Klassifizierung zu 0,1 kommt durch die Wahl des AU-Modells zustande. Dabei wird SVM gewählt. (ADABase Paper)\n",
|
||||
"Gibt die Klassifizierung einen Gleitkommawert zwischen 0 & 1 aus, dann kommt XGB zum Einsatz. (REVELIO Paper)"
|
||||
"### Python Environment Configuration\n",
|
||||
"\n",
|
||||
"**Py-Feat** relies on dependencies that are incompatible with Python 3.12 and later. To ensure functionality, you must set up a dedicated **Python 3.11** kernel.\n",
|
||||
"\n",
|
||||
"#### Setup Instructions:\n",
|
||||
"\n",
|
||||
"1. Open your **Jupyter Lab** interface.\n",
|
||||
"2. Open a **Terminal** and execute the following commands:\n",
|
||||
"```bash\n",
|
||||
"conda create -n py311 python=3.11\n",
|
||||
"source ~/.bashrc\n",
|
||||
"conda activate py311\n",
|
||||
"conda install jupyter\n",
|
||||
"python -m ipykernel install --user --name=py311 --display-name \"Python 3.11\"\n",
|
||||
"pip install py-feat\n",
|
||||
"pip install \"moviepy<2.0\" # Only if required\n",
|
||||
"\n",
|
||||
"```\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"3. **Restart** the kernel.\n",
|
||||
"4. In **VS Code**, refresh your kernel list and select the one labeled **\"Python 3.11\"**.\n",
|
||||
"\n",
|
||||
"---\n",
|
||||
"\n",
|
||||
"### Implementation Details\n",
|
||||
"\n",
|
||||
"The following code demonstrates a sample integration of the `py-feat` library. The classification output format is determined by the specific AU model selected:\n",
|
||||
"\n",
|
||||
"| Model | Output Type | Reference Paper |\n",
|
||||
"| --- | --- | --- |\n",
|
||||
"| **SVM** | Binary (0 or 1) | *ADABase* |\n",
|
||||
"| **XGB** | Floating Point (0.0 - 1.0) | *REVELIO* |\n",
|
||||
"\n",
|
||||
"---\n",
|
||||
"\n",
|
||||
"Would you like me to provide the Python code block to implement the **SVM** or **XGB** detector using these libraries?"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
import cv2
|
||||
import time
|
||||
import os
|
||||
import threading
|
||||
from datetime import datetime
|
||||
from feat import Detector
|
||||
import torch
|
||||
import pandas as pd
|
||||
|
||||
# Import your helper functions
|
||||
# from db_helper import connect_db, disconnect_db, insert_rows_into_table, create_table
|
||||
import db_helper as db
|
||||
|
||||
|
||||
# Konfiguration
|
||||
DB_PATH = "action_units.db" # TODO
|
||||
CAMERA_INDEX = 0
|
||||
OUTPUT_DIR = "recordings"
|
||||
VIDEO_DURATION = 50 # Sekunden
|
||||
START_INTERVAL = 5 # Sekunden bis zum nächsten Start
|
||||
FPS = 25.0 # Feste FPS
|
||||
|
||||
if not os.path.exists(OUTPUT_DIR):
|
||||
os.makedirs(OUTPUT_DIR)
|
||||
|
||||
# Globaler Detector, um ihn nicht bei jedem Video neu laden zu müssen (spart massiv Zeit/Speicher)
|
||||
print("Initialisiere AU-Detector (bitte warten)...")
|
||||
detector = Detector(au_model="xgb")
|
||||
|
||||
def extract_aus(path, skip_frames):
|
||||
|
||||
# torch.no_grad() deaktiviert die Gradientenberechnung.
|
||||
# Das löst den "Can't call numpy() on Tensor that requires grad" Fehler.
|
||||
with torch.no_grad():
|
||||
video_prediction = detector.detect_video(
|
||||
path,
|
||||
skip_frames=skip_frames,
|
||||
face_detection_threshold=0.95
|
||||
)
|
||||
|
||||
# Falls video_prediction oder .aus noch Tensoren sind,
|
||||
# stellen wir sicher, dass sie korrekt summiert werden.
|
||||
try:
|
||||
# Wir nehmen die Summe der Action Units über alle detektierten Frames
|
||||
res = video_prediction.aus.sum()
|
||||
return res
|
||||
except Exception as e:
|
||||
print(f"Fehler bei der Summenbildung: {e}")
|
||||
return None
|
||||
|
||||
def startAU_creation(video_path, db_path):
|
||||
"""Diese Funktion läuft nun in einem eigenen Thread."""
|
||||
try:
|
||||
print(f"\n[THREAD START] Analyse läuft für: {video_path}")
|
||||
# skip_frames berechnen (z.B. alle 5 Sekunden bei 25 FPS = 125)
|
||||
output = extract_aus(video_path, skip_frames=int(FPS*5))
|
||||
|
||||
print(f"\n--- Ergebnis für {os.path.basename(video_path)} ---")
|
||||
print(output)
|
||||
print("--------------------------------------------------\n")
|
||||
if output is not None:
|
||||
# Verbindung für diesen Thread öffnen (SQLite Sicherheit)
|
||||
conn, cursor = db.connect_db(db_path)
|
||||
|
||||
# Daten vorbereiten: Timestamp + AU Ergebnisse
|
||||
# Wir wandeln die Series/Dataframe in ein Dictionary um
|
||||
data_to_insert = output.to_dict()
|
||||
data_to_insert['timestamp'] = [datetime.now().strftime("%Y-%m-%d %H:%M:%S")]
|
||||
|
||||
# Da die AU-Spaltennamen dynamisch sind, stellen wir sicher, dass sie Listen sind
|
||||
# (insert_rows_into_table erwartet Listen für jeden Key)
|
||||
final_payload = {k: [v] if not isinstance(v, list) else v for k, v in data_to_insert.items()}
|
||||
|
||||
|
||||
db.insert_rows_into_table(conn, cursor, "actionUnits", final_payload)
|
||||
|
||||
db.disconnect_db(conn, cursor)
|
||||
print(f"--- Ergebnis für {os.path.basename(video_path)} in DB gespeichert ---")
|
||||
except Exception as e:
|
||||
print(f"Fehler bei der Analyse von {video_path}: {e}")
|
||||
|
||||
class VideoRecorder:
|
||||
def __init__(self, filename, width, height, db_path):
|
||||
self.filename = filename
|
||||
self.db_path = db_path
|
||||
fourcc = cv2.VideoWriter_fourcc(*'XVID')
|
||||
self.out = cv2.VideoWriter(filename, fourcc, FPS, (width, height))
|
||||
self.frames_to_record = int(VIDEO_DURATION * FPS)
|
||||
self.frames_count = 0
|
||||
self.is_finished = False
|
||||
|
||||
def write_frame(self, frame):
|
||||
if self.frames_count < self.frames_to_record:
|
||||
self.out.write(frame)
|
||||
self.frames_count += 1
|
||||
else:
|
||||
self.finish()
|
||||
|
||||
def finish(self):
|
||||
if not self.is_finished:
|
||||
self.out.release()
|
||||
self.is_finished = True
|
||||
abs_path = os.path.abspath(self.filename)
|
||||
print(f"Video fertig gespeichert: {self.filename}")
|
||||
|
||||
# --- MULTITHREADING HIER ---
|
||||
# Wir starten die Analyse in einem neuen Thread, damit main() sofort weiter frames lesen kann
|
||||
analysis_thread = threading.Thread(target=startAU_creation, args=(abs_path, self.db_path))
|
||||
analysis_thread.daemon = True # Beendet sich, wenn das Hauptprogramm schließt
|
||||
analysis_thread.start()
|
||||
|
||||
def main():
|
||||
cap = cv2.VideoCapture(CAMERA_INDEX)
|
||||
if not cap.isOpened():
|
||||
print("Fehler: Kamera konnte nicht geöffnet werden.")
|
||||
return
|
||||
|
||||
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
active_recorders = []
|
||||
last_start_time = 0
|
||||
|
||||
print("Aufnahme läuft. Drücke 'q' zum Beenden.")
|
||||
|
||||
try:
|
||||
while True:
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
|
||||
current_time = time.time()
|
||||
|
||||
if current_time - last_start_time >= START_INTERVAL:
|
||||
timestamp = datetime.now().strftime("%H%M%S")
|
||||
filename = os.path.join(OUTPUT_DIR, f"rec_{timestamp}.avi")
|
||||
new_recorder = VideoRecorder(filename, width, height, DB_PATH)
|
||||
active_recorders.append(new_recorder)
|
||||
last_start_time = current_time
|
||||
|
||||
for rec in active_recorders[:]:
|
||||
rec.write_frame(frame)
|
||||
if rec.is_finished:
|
||||
active_recorders.remove(rec)
|
||||
|
||||
cv2.imshow('Kamera Livestream', frame)
|
||||
if cv2.waitKey(1) & 0xFF == ord('q'):
|
||||
break
|
||||
|
||||
time.sleep(1/FPS)
|
||||
|
||||
finally:
|
||||
cap.release()
|
||||
cv2.destroyAllWindows()
|
||||
print("Programm beendet. Warte ggf. auf laufende Analysen...")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,371 @@
|
||||
import cv2
|
||||
import time
|
||||
import os
|
||||
import threading
|
||||
import warnings
|
||||
from datetime import datetime
|
||||
from feat import Detector
|
||||
import torch
|
||||
import mediapipe as mp
|
||||
import pandas as pd
|
||||
import db_helper as db
|
||||
|
||||
from pathlib import Path
|
||||
from eyeFeature_new import compute_features_from_parquet
|
||||
|
||||
# Suppress specific Protobuf deprecation warnings from the library
|
||||
warnings.filterwarnings(
|
||||
"ignore",
|
||||
message=r".*SymbolDatabase\.GetPrototype\(\) is deprecated.*",
|
||||
category=UserWarning,
|
||||
module=r"google\.protobuf\.symbol_database"
|
||||
)
|
||||
|
||||
# --- Configuration & Hyperparameters ---
|
||||
DB_PATH = Path("~/MSY_FS/databases/database.sqlite").expanduser()
|
||||
CAMERA_INDEX = 0
|
||||
OUTPUT_DIR = Path("recordings")
|
||||
VIDEO_DURATION = 50 # Seconds per recording segment
|
||||
START_INTERVAL = 5 # Delay between starting overlapping recordings
|
||||
FPS = 25.0 # Target Frames Per Second
|
||||
|
||||
# Global feature storage - Updated to be thread-safe in production environments
|
||||
eye_tracking_features = {}
|
||||
|
||||
if not OUTPUT_DIR.exists():
|
||||
OUTPUT_DIR.mkdir(parents=True)
|
||||
|
||||
# Initialize the AU-Detector globally to optimize VRAM/RAM usage
|
||||
print("[INFO] Initializing Facial Action Unit Detector (XGB)...")
|
||||
detector = Detector(au_model="xgb")
|
||||
|
||||
# --- MediaPipe FaceMesh Configuration ---
|
||||
mp_face_mesh = mp.solutions.face_mesh
|
||||
face_mesh = mp_face_mesh.FaceMesh(
|
||||
static_image_mode=False,
|
||||
max_num_faces=1,
|
||||
refine_landmarks=True, # Mandatory for Iris tracking
|
||||
min_detection_confidence=0.5,
|
||||
min_tracking_confidence=0.5
|
||||
)
|
||||
|
||||
# Landmark Indices for Oculometrics
|
||||
LEFT_IRIS = [474, 475, 476, 477]
|
||||
RIGHT_IRIS = [469, 470, 471, 472]
|
||||
LEFT_EYE_LIDS = (159, 145)
|
||||
RIGHT_EYE_LIDS = (386, 374)
|
||||
EYE_OPEN_THRESHOLD = 6
|
||||
|
||||
# Bounding box indices for eye regions
|
||||
LEFT_EYE_ALL = [33, 7, 163, 144, 145, 153, 154, 155, 133, 173, 157, 158, 159, 160, 161, 246]
|
||||
RIGHT_EYE_ALL = [263, 249, 390, 373, 374, 380, 381, 382, 362, 398, 384, 385, 386, 387, 388, 466]
|
||||
|
||||
def eye_openness(landmarks, top_idx, bottom_idx, img_height):
|
||||
"""Calculates the vertical distance between eyelids normalized by image height."""
|
||||
top = landmarks[top_idx]
|
||||
bottom = landmarks[bottom_idx]
|
||||
return abs(top.y - bottom.y) * img_height
|
||||
|
||||
def compute_gaze(landmarks, iris_center, eye_indices, w, h):
|
||||
"""
|
||||
Computes normalized gaze coordinates (0.0 to 1.0) relative to the eye's
|
||||
internal bounding box.
|
||||
"""
|
||||
iris_x, iris_y = iris_center
|
||||
|
||||
eye_points = []
|
||||
for idx in eye_indices:
|
||||
lm = landmarks[idx]
|
||||
eye_points.append((lm.x * w, lm.y * h))
|
||||
|
||||
xs = [p[0] for p in eye_points]
|
||||
ys = [p[1] for p in eye_points]
|
||||
|
||||
eye_left = min(xs)
|
||||
eye_right = max(xs)
|
||||
eye_top = min(ys)
|
||||
eye_bottom = max(ys)
|
||||
|
||||
eye_width = eye_right - eye_left
|
||||
eye_height = eye_bottom - eye_top
|
||||
|
||||
if eye_width < 1 or eye_height < 1:
|
||||
return 0.5, 0.5
|
||||
|
||||
gaze_x = (iris_x - eye_left) / eye_width
|
||||
gaze_y = (iris_y - eye_top) / eye_height
|
||||
|
||||
return gaze_x, gaze_y
|
||||
|
||||
def extract_aus(path, skip_frames):
|
||||
"""
|
||||
Infers facial Action Units from video file.
|
||||
Uses torch.no_grad() to optimize inference and prevent memory leakage.
|
||||
"""
|
||||
with torch.no_grad():
|
||||
try:
|
||||
video_prediction = detector.detect_video(
|
||||
path,
|
||||
skip_frames=skip_frames,
|
||||
face_detection_threshold=0.95
|
||||
)
|
||||
# Compute temporal mean of Action Units across the segment
|
||||
return video_prediction.aus.mean()
|
||||
except Exception as e:
|
||||
print(f"[ERROR] AU Extraction failed: {e}")
|
||||
return None
|
||||
|
||||
def process_and_store_analysis(video_path, db_path):
|
||||
"""
|
||||
Worker function: Handles AU extraction, data merging, and SQL persistence.
|
||||
Designed to run in a background thread.
|
||||
"""
|
||||
try:
|
||||
print(f"[THREAD] Analyzing segment: {video_path}")
|
||||
# Analysis sampling: one frame every 5 seconds
|
||||
output = extract_aus(video_path, skip_frames=int(FPS * 5))
|
||||
if output is not None:
|
||||
# Verbindung für diesen Thread öffnen (SQLite Sicherheit)
|
||||
conn, cursor = db.connect_db(db_path)
|
||||
|
||||
# Prepare payload: Prefix keys to distinguish facial AUs
|
||||
data_to_insert = output.to_dict()
|
||||
|
||||
data_to_insert = {
|
||||
f"FACE_{k}_mean": v for k, v in data_to_insert.items()
|
||||
}
|
||||
|
||||
now = datetime.now()
|
||||
ticks = int(time.mktime(now.timetuple()))
|
||||
|
||||
data_to_insert['start_time'] = [ticks]
|
||||
data_to_insert = data_to_insert | eye_tracking_features
|
||||
|
||||
# making sure that dynamic AU-columns are lists
|
||||
# (insert_rows_into_table expects lists for every key)
|
||||
final_payload = {k: [v] if not isinstance(v, list) else v for k, v in data_to_insert.items()}
|
||||
|
||||
|
||||
db.insert_rows_into_table(conn, cursor, "feature_table", final_payload)
|
||||
|
||||
db.disconnect_db(conn, cursor)
|
||||
print(f"[SUCCESS] Data persisted for {os.path.basename(video_path)}")
|
||||
|
||||
# Cleanup temporary files to save disk space
|
||||
os.remove(video_path)
|
||||
os.remove(video_path.replace(".avi", "_gaze.parquet"))
|
||||
|
||||
|
||||
except Exception as e:
|
||||
print(f"[ERROR] Threaded analysis failed for {video_path}: {e}")
|
||||
|
||||
|
||||
class VideoRecorder:
|
||||
"""Manages the asynchronous writing of video frames to disk."""
|
||||
def __init__(self, filename, width, height, db_path):
|
||||
self.gaze_data = []
|
||||
self.filename = filename
|
||||
self.db_path = db_path
|
||||
fourcc = cv2.VideoWriter_fourcc(*'XVID')
|
||||
self.out = cv2.VideoWriter(filename, fourcc, FPS, (width, height))
|
||||
self.frames_to_record = int(VIDEO_DURATION * FPS)
|
||||
self.frames_count = 0
|
||||
self.is_finished = False
|
||||
|
||||
def write_frame(self, frame):
|
||||
if self.frames_count < self.frames_to_record:
|
||||
self.out.write(frame)
|
||||
self.frames_count += 1
|
||||
else:
|
||||
self.finish()
|
||||
|
||||
def finish(self):
|
||||
if not self.is_finished:
|
||||
self.out.release()
|
||||
self.is_finished = True
|
||||
abs_path = os.path.abspath(self.filename)
|
||||
print(f"Video saved: {self.filename}")
|
||||
|
||||
# Trigger background analysis thread
|
||||
# Passing a snapshot of eye_tracking_features to avoid race conditions
|
||||
analysis_thread = threading.Thread(target=process_and_store_analysis, args=(abs_path, self.db_path))
|
||||
analysis_thread.daemon = True # ends when the program ends
|
||||
analysis_thread.start()
|
||||
|
||||
class GazeRecorder:
|
||||
"""Handles the collection and Parquet serialization of oculometric data."""
|
||||
def __init__(self, filename):
|
||||
self.filename = filename
|
||||
self.frames_to_record = int(VIDEO_DURATION * FPS)
|
||||
self.frames_count = 0
|
||||
self.gaze_data = []
|
||||
self.is_finished = False
|
||||
|
||||
def write_frame(self, gaze_row):
|
||||
if self.frames_count < self.frames_to_record:
|
||||
self.gaze_data.append(gaze_row)
|
||||
self.frames_count += 1
|
||||
else:
|
||||
self.finish()
|
||||
|
||||
def finish(self):
|
||||
if not self.is_finished:
|
||||
df = pd.DataFrame(self.gaze_data)
|
||||
df.to_parquet(self.filename, engine="pyarrow", index=False)
|
||||
|
||||
# Extract high-level features from raw gaze points
|
||||
print(f"Gaze-Parquet saved: {self.filename}")
|
||||
features = compute_features_from_parquet(self.filename)
|
||||
print("Features:", features)
|
||||
self.is_finished = True
|
||||
eye_tracking_features = features
|
||||
|
||||
def main():
|
||||
cap = cv2.VideoCapture(CAMERA_INDEX)
|
||||
if not cap.isOpened():
|
||||
print("[CRITICAL] Could not access camera.")
|
||||
return
|
||||
|
||||
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
|
||||
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
|
||||
|
||||
active_video_recorders = []
|
||||
active_gaze_recorders = []
|
||||
last_start_time = 0
|
||||
|
||||
print("[INFO] Recording started. Press 'q' to terminate.")
|
||||
|
||||
try:
|
||||
while True:
|
||||
ret, frame = cap.read()
|
||||
if not ret:
|
||||
break
|
||||
|
||||
# Pre-processing for MediaPipe
|
||||
rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
h, w, _ = frame.shape
|
||||
results = face_mesh.process(rgb)
|
||||
|
||||
# Default feature values
|
||||
left_valid = 0
|
||||
right_valid = 0
|
||||
left_diameter = None
|
||||
right_diameter = None
|
||||
|
||||
left_gaze_x = None
|
||||
left_gaze_y = None
|
||||
right_gaze_x = None
|
||||
right_gaze_y = None
|
||||
|
||||
if results.multi_face_landmarks:
|
||||
face_landmarks = results.multi_face_landmarks[0]
|
||||
|
||||
left_open = eye_openness(
|
||||
face_landmarks.landmark,
|
||||
LEFT_EYE_LIDS[0],
|
||||
LEFT_EYE_LIDS[1],
|
||||
h
|
||||
)
|
||||
|
||||
right_open = eye_openness(
|
||||
face_landmarks.landmark,
|
||||
RIGHT_EYE_LIDS[0],
|
||||
RIGHT_EYE_LIDS[1],
|
||||
h
|
||||
)
|
||||
|
||||
left_valid = 1 if left_open > EYE_OPEN_THRESHOLD else 0
|
||||
right_valid = 1 if right_open > EYE_OPEN_THRESHOLD else 0
|
||||
|
||||
for eye_name, eye_indices in [("left", LEFT_IRIS), ("right", RIGHT_IRIS)]:
|
||||
iris_points = []
|
||||
|
||||
for idx in eye_indices:
|
||||
lm = face_landmarks.landmark[idx]
|
||||
x_i, y_i = int(lm.x * w), int(lm.y * h)
|
||||
iris_points.append((x_i, y_i))
|
||||
|
||||
if len(iris_points) == 4:
|
||||
cx = int(sum(p[0] for p in iris_points) / 4)
|
||||
cy = int(sum(p[1] for p in iris_points) / 4)
|
||||
|
||||
radius = max(
|
||||
((x - cx) ** 2 + (y - cy) ** 2) ** 0.5
|
||||
for (x, y) in iris_points
|
||||
)
|
||||
|
||||
diameter = 2 * radius
|
||||
|
||||
cv2.circle(frame, (cx, cy), int(radius), (0, 255, 0), 2)
|
||||
|
||||
if eye_name == "left" and left_valid:
|
||||
left_diameter = diameter
|
||||
left_gaze_x, left_gaze_y = compute_gaze(
|
||||
face_landmarks.landmark,
|
||||
(cx, cy),
|
||||
RIGHT_EYE_ALL,
|
||||
w, h
|
||||
)
|
||||
|
||||
elif eye_name == "right" and right_valid:
|
||||
right_diameter = diameter
|
||||
right_gaze_x, right_gaze_y = compute_gaze(
|
||||
face_landmarks.landmark,
|
||||
(cx, cy),
|
||||
LEFT_EYE_ALL,
|
||||
w, h
|
||||
)
|
||||
|
||||
gaze_row = {
|
||||
"timestamp": time.time(),
|
||||
"EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_X": left_gaze_x,
|
||||
"EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_Y": left_gaze_y,
|
||||
"EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_X": right_gaze_x,
|
||||
"EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_Y": right_gaze_y,
|
||||
"EYE_LEFT_PUPIL_VALIDITY": left_valid,
|
||||
"EYE_RIGHT_PUPIL_VALIDITY": right_valid,
|
||||
"EYE_LEFT_PUPIL_DIAMETER": left_diameter,
|
||||
"EYE_RIGHT_PUPIL_DIAMETER": right_diameter
|
||||
}
|
||||
|
||||
current_time = time.time()
|
||||
|
||||
if current_time - last_start_time >= START_INTERVAL:
|
||||
timestamp = datetime.now().strftime("%H%M%S")
|
||||
filename = os.path.join(OUTPUT_DIR, f"rec_{timestamp}.avi")
|
||||
video_recorder = VideoRecorder(filename, width, height, DB_PATH)
|
||||
|
||||
gaze_filename = filename.replace(".avi", "_gaze.parquet")
|
||||
gaze_recorder = GazeRecorder(gaze_filename)
|
||||
|
||||
active_video_recorders.append(video_recorder)
|
||||
active_gaze_recorders.append(gaze_recorder)
|
||||
|
||||
last_start_time = current_time
|
||||
|
||||
for v_rec, g_rec in zip(active_video_recorders[:], active_gaze_recorders[:]):
|
||||
|
||||
v_rec.write_frame(frame)
|
||||
g_rec.write_frame(gaze_row)
|
||||
|
||||
if v_rec.is_finished:
|
||||
active_video_recorders.remove(v_rec)
|
||||
|
||||
if g_rec.is_finished:
|
||||
active_gaze_recorders.remove(g_rec)
|
||||
|
||||
cv2.imshow('Kamera Livestream', frame)
|
||||
if cv2.waitKey(1) & 0xFF == ord('q'):
|
||||
break
|
||||
|
||||
time.sleep(1/FPS)
|
||||
|
||||
finally:
|
||||
face_mesh.close()
|
||||
cap.release()
|
||||
cv2.destroyAllWindows()
|
||||
|
||||
print("[INFO] Stream closed. Waiting for background analysis to complete...")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,166 @@
|
||||
import os
|
||||
import sqlite3
|
||||
import pandas as pd
|
||||
|
||||
|
||||
def connect_db(path_to_file: os.PathLike) -> tuple[sqlite3.Connection, sqlite3.Cursor]:
|
||||
''' Establishes a connection with a sqlite3 database. '''
|
||||
conn = sqlite3.connect(path_to_file)
|
||||
cursor = conn.cursor()
|
||||
return conn, cursor
|
||||
|
||||
def disconnect_db(conn: sqlite3.Connection, cursor: sqlite3.Cursor, commit: bool = True) -> None:
|
||||
''' Commits all remaining changes and closes the connection with an sqlite3 database. '''
|
||||
cursor.close()
|
||||
if commit: conn.commit() # commit all pending changes made to the sqlite3 database before closing
|
||||
conn.close()
|
||||
|
||||
def create_table(
|
||||
conn: sqlite3.Connection,
|
||||
cursor: sqlite3.Cursor,
|
||||
table_name: str,
|
||||
columns: dict,
|
||||
constraints: dict,
|
||||
primary_key: dict,
|
||||
commit: bool = True
|
||||
) -> str:
|
||||
'''
|
||||
Creates a new empty table with the given columns, constraints and primary key.
|
||||
|
||||
:param columns: dict with column names (=keys) and dtypes (=values) (e.g. BIGINT, INT, ...)
|
||||
:param constraints: dict with column names (=keys) and list of constraints (=values) (like [\'NOT NULL\'(,...)])
|
||||
:param primary_key: dict with primary key name (=key) and list of attributes which combined define the table's primary key (=values, like [\'att1\'(,...)])
|
||||
'''
|
||||
assert len(primary_key.keys()) == 1
|
||||
sql = f'CREATE TABLE {table_name} (\n '
|
||||
for column,dtype in columns.items():
|
||||
sql += f'{column} {dtype}{" "+" ".join(constraints[column]) if column in constraints.keys() else ""},\n '
|
||||
if list(primary_key.keys())[0]: sql += f'CONSTRAINT {list(primary_key.keys())[0]} '
|
||||
sql += f'PRIMARY KEY ({", ".join(list(primary_key.values())[0])})\n)'
|
||||
cursor.execute(sql)
|
||||
if commit: conn.commit()
|
||||
return sql
|
||||
|
||||
def add_columns_to_table(
|
||||
conn: sqlite3.Connection,
|
||||
cursor: sqlite3.Cursor,
|
||||
table_name: str,
|
||||
columns: dict,
|
||||
constraints: dict = dict(),
|
||||
commit: bool = True
|
||||
) -> str:
|
||||
''' Adds one/multiple columns (each with a list of constraints) to the given table. '''
|
||||
sql_total = ''
|
||||
for column,dtype in columns.items(): # sqlite can only add one column per query
|
||||
sql = f'ALTER TABLE {table_name}\n '
|
||||
sql += f'ADD "{column}" {dtype}{" "+" ".join(constraints[column]) if column in constraints.keys() else ""}'
|
||||
sql_total += sql + '\n'
|
||||
cursor.execute(sql)
|
||||
if commit: conn.commit()
|
||||
return sql_total
|
||||
|
||||
|
||||
|
||||
|
||||
def insert_rows_into_table(
|
||||
conn: sqlite3.Connection,
|
||||
cursor: sqlite3.Cursor,
|
||||
table_name: str,
|
||||
columns: dict,
|
||||
commit: bool = True
|
||||
) -> str:
|
||||
'''
|
||||
Inserts values as multiple rows into the given table.
|
||||
|
||||
:param columns: dict with column names (=keys) and values to insert as lists with at least one element (=values)
|
||||
|
||||
Note: The number of given values per attribute must match the number of rows to insert!
|
||||
Note: The values for the rows must be of normal python types (e.g. list, str, int, ...) instead of e.g. numpy arrays!
|
||||
'''
|
||||
assert len(set(map(len, columns.values()))) == 1, 'ERROR: Provide equal number of values for each column!'
|
||||
assert len(set(list(map(type,columns.values())))) == 1 and isinstance(list(columns.values())[0], list), 'ERROR: Provide values as Python lists!'
|
||||
assert set([type(a) for b in list(columns.values()) for a in b]).issubset({str,int,float,bool}), 'ERROR: Provide values as basic Python data types!'
|
||||
|
||||
values = list(zip(*columns.values()))
|
||||
sql = f'INSERT INTO {table_name} ({", ".join(columns.keys())})\n VALUES ({("?,"*len(values[0]))[:-1]})'
|
||||
cursor.executemany(sql, values)
|
||||
if commit: conn.commit()
|
||||
return sql
|
||||
|
||||
def update_multiple_rows_in_table(
|
||||
conn: sqlite3.Connection,
|
||||
cursor: sqlite3.Cursor,
|
||||
table_name: str,
|
||||
new_vals: dict,
|
||||
conditions: str,
|
||||
commit: bool = True
|
||||
) -> str:
|
||||
'''
|
||||
Updates attribute values of some rows in the given table.
|
||||
|
||||
:param new_vals: dict with column names (=keys) and the new values to set (=values)
|
||||
:param conditions: string which defines all concatenated conditions (e.g. \'cond1 AND (cond2 OR cond3)\' with cond1: att1=5, ...)
|
||||
'''
|
||||
assignments = ', '.join([f'{k}={v}' for k,v in zip(new_vals.keys(), new_vals.values())])
|
||||
sql = f'UPDATE {table_name}\n SET {assignments}\n WHERE {conditions}'
|
||||
cursor.execute(sql)
|
||||
if commit: conn.commit()
|
||||
return sql
|
||||
|
||||
def delete_rows_from_table(
|
||||
conn: sqlite3.Connection,
|
||||
cursor: sqlite3.Cursor,
|
||||
table_name: str,
|
||||
conditions: str,
|
||||
commit: bool = True
|
||||
) -> str:
|
||||
'''
|
||||
Deletes rows from the given table.
|
||||
|
||||
:param conditions: string which defines all concatenated conditions (e.g. \'cond1 AND (cond2 OR cond3)\' with cond1: att1=5, ...)
|
||||
'''
|
||||
sql = f'DELETE FROM {table_name} WHERE {conditions}'
|
||||
cursor.execute(sql)
|
||||
if commit: conn.commit()
|
||||
return sql
|
||||
|
||||
|
||||
|
||||
def get_data_from_table(
|
||||
conn: sqlite3.Connection,
|
||||
table_name: str,
|
||||
columns_list: list = ['*'],
|
||||
aggregations: [None,dict] = None,
|
||||
where_conditions: [None,str] = None,
|
||||
order_by: [None, dict] = None,
|
||||
limit: [None, int] = None,
|
||||
offset: [None, int] = None
|
||||
) -> pd.DataFrame:
|
||||
'''
|
||||
Helper function which returns (if desired: aggregated) contents from the given table as a pandas DataFrame. The rows can be filtered by providing the condition as a string.
|
||||
|
||||
:param columns_list: use if no aggregation is needed to select which columns to get from the table
|
||||
:param (optional) aggregations: use to apply aggregations on the data from the table; dictionary with column(s) as key(s) and aggregation(s) as corresponding value(s) (e.g. {'col1': 'MIN', 'col2': 'AVG', ...} or {'*': 'COUNT'})
|
||||
:param (optional) where_conditions: string which defines all concatenated conditions (e.g. \'cond1 AND (cond2 OR cond3)\' with cond1: att1=5, ...) applied on table.
|
||||
:param (optional) order_by: dict defining the ordering of the outputs with column(s) as key(s) and ordering as corresponding value(s) (e.g. {'col1': 'ASC'})
|
||||
:param (optional) limit: use to limit the number of returned rows
|
||||
:param (optional) offset: use to skip the first n rows before displaying
|
||||
|
||||
Note: If aggregations is set, the columns_list is ignored.
|
||||
Note: Get all data as a DataFrame with get_data_from_table(conn, table_name).
|
||||
Note: If one output is wanted (e.g. count(*) or similar), get it with get_data_from_table(...).iloc[0,0] from the DataFrame.
|
||||
'''
|
||||
assert columns_list or aggregations
|
||||
|
||||
if aggregations:
|
||||
selection = [f'{agg}({col})' for col,agg in aggregations.items()]
|
||||
else:
|
||||
selection = columns_list
|
||||
selection = ", ".join(selection)
|
||||
where_conditions = 'WHERE ' + where_conditions if where_conditions else ''
|
||||
order_by = 'ORDER BY ' + ', '.join([f'{k} {v}' for k,v in order_by.items()]) if order_by else ''
|
||||
limit = f'LIMIT {limit}' if limit else ''
|
||||
offset = f'OFFSET {offset}' if offset else ''
|
||||
|
||||
sql = f'SELECT {selection} FROM {table_name} {where_conditions} {order_by} {limit} {offset}'
|
||||
return pd.read_sql_query(sql, conn)
|
||||
@@ -0,0 +1,54 @@
|
||||
import db_helper as db
|
||||
|
||||
DB_PATH = "action_units.db"
|
||||
|
||||
def setup_test_db():
|
||||
# 1. Verbindung herstellen (erstellt die Datei, falls nicht vorhanden)
|
||||
conn, cursor = db.connect_db(DB_PATH)
|
||||
|
||||
# 2. Spalten definieren
|
||||
# Wir erstellen eine Spalte für den Zeitstempel und beispielhaft einige AUs.
|
||||
# In SQLite können wir später mit deinem Helper weitere Spalten hinzufügen.
|
||||
columns = {
|
||||
"timestamp": "TEXT",
|
||||
"AU01": "REAL",
|
||||
"AU02": "REAL",
|
||||
"AU04": "REAL",
|
||||
"AU05": "REAL",
|
||||
"AU06": "REAL",
|
||||
"AU07": "REAL",
|
||||
"AU09": "REAL",
|
||||
"AU10": "REAL",
|
||||
"AU11": "REAL",
|
||||
"AU12": "REAL",
|
||||
"AU14": "REAL",
|
||||
"AU15": "REAL",
|
||||
"AU17": "REAL",
|
||||
"AU20": "REAL",
|
||||
"AU23": "REAL",
|
||||
"AU24": "REAL",
|
||||
"AU25": "REAL",
|
||||
"AU26": "REAL",
|
||||
"AU28": "REAL",
|
||||
"AU43": "REAL",
|
||||
}
|
||||
|
||||
# Constraints (z.B. Zeitstempel darf nicht leer sein)
|
||||
constraints = {
|
||||
"timestamp": ["NOT NULL"]
|
||||
}
|
||||
|
||||
# Primärschlüssel definieren (Kombination aus Zeitstempel und ggf. ID)
|
||||
primary_key = {"pk_timestamp": ["timestamp"]}
|
||||
|
||||
try:
|
||||
sql = db.create_table(conn, cursor, "actionUnits", columns, constraints, primary_key)
|
||||
print("Tabelle erfolgreich erstellt!")
|
||||
print(f"SQL-Befehl:\n{sql}")
|
||||
except Exception as e:
|
||||
print(f"Hinweis: {e}")
|
||||
finally:
|
||||
db.disconnect_db(conn, cursor)
|
||||
|
||||
if __name__ == "__main__":
|
||||
setup_test_db()
|
||||
@@ -0,0 +1,174 @@
|
||||
import cv2
|
||||
import mediapipe as mp
|
||||
import numpy as np
|
||||
import pyautogui
|
||||
import pandas as pd
|
||||
import time
|
||||
from sklearn.pipeline import make_pipeline
|
||||
from sklearn.preprocessing import PolynomialFeatures
|
||||
from sklearn.linear_model import LinearRegression
|
||||
|
||||
# Bildschirmgröße
|
||||
screen_w, screen_h = pyautogui.size()
|
||||
|
||||
# MediaPipe Setup
|
||||
mp_face_mesh = mp.solutions.face_mesh
|
||||
face_mesh = mp_face_mesh.FaceMesh(refine_landmarks=True)
|
||||
cap = cv2.VideoCapture(0)
|
||||
|
||||
# Iris Landmark Indizes
|
||||
LEFT_IRIS = [468, 469, 470, 471, 472]
|
||||
RIGHT_IRIS = [473, 474, 475, 476, 477]
|
||||
|
||||
def get_iris_center(landmarks, indices):
|
||||
points = np.array([[landmarks[i].x, landmarks[i].y] for i in indices])
|
||||
return np.mean(points, axis=0)
|
||||
|
||||
# Kalibrierpunkte
|
||||
calibration_points = [
|
||||
(0.1,0.1),(0.5,0.1),(0.9,0.1),
|
||||
(0.1,0.5),(0.5,0.5),(0.9,0.5),
|
||||
(0.1,0.9),(0.5,0.9),(0.9,0.9)
|
||||
]
|
||||
|
||||
left_data = []
|
||||
right_data = []
|
||||
|
||||
print("Kalibrierung startet...")
|
||||
|
||||
for idx, (px, py) in enumerate(calibration_points):
|
||||
|
||||
screen = np.zeros((screen_h, screen_w, 3), dtype=np.uint8)
|
||||
|
||||
for j, (cpx, cpy) in enumerate(calibration_points):
|
||||
cx = int(cpx * screen_w)
|
||||
cy = int(cpy * screen_h)
|
||||
|
||||
if j == idx:
|
||||
color = (0, 0, 255)
|
||||
radius = 25
|
||||
else:
|
||||
color = (255, 255, 255)
|
||||
radius = 15
|
||||
|
||||
cv2.circle(screen, (cx, cy), radius, color, -1)
|
||||
|
||||
# Fenster vorbereiten
|
||||
cv2.namedWindow("Calibration", cv2.WINDOW_NORMAL)
|
||||
cv2.imshow("Calibration", screen)
|
||||
cv2.waitKey(1000)
|
||||
|
||||
samples_left = []
|
||||
samples_right = []
|
||||
|
||||
start = time.time()
|
||||
while time.time() - start < 2:
|
||||
ret, frame = cap.read()
|
||||
frame = cv2.flip(frame, 1)
|
||||
rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
results = face_mesh.process(rgb)
|
||||
|
||||
if results.multi_face_landmarks:
|
||||
mesh = results.multi_face_landmarks[0].landmark
|
||||
|
||||
left_center = get_iris_center(mesh, LEFT_IRIS)
|
||||
right_center = get_iris_center(mesh, RIGHT_IRIS)
|
||||
|
||||
samples_left.append(left_center)
|
||||
samples_right.append(right_center)
|
||||
|
||||
avg_left = np.mean(samples_left, axis=0)
|
||||
avg_right = np.mean(samples_right, axis=0)
|
||||
|
||||
target_x = int(px * screen_w)
|
||||
target_y = int(py * screen_h)
|
||||
|
||||
left_data.append([avg_left[0], avg_left[1], target_x, target_y])
|
||||
right_data.append([avg_right[0], avg_right[1], target_x, target_y])
|
||||
|
||||
cv2.destroyWindow("Calibration")
|
||||
|
||||
# Training
|
||||
def train_model(data):
|
||||
data = np.array(data)
|
||||
X = data[:, :2]
|
||||
yx = data[:, 2]
|
||||
yy = data[:, 3]
|
||||
|
||||
model_x = make_pipeline(PolynomialFeatures(2), LinearRegression())
|
||||
model_y = make_pipeline(PolynomialFeatures(2), LinearRegression())
|
||||
|
||||
model_x.fit(X, yx)
|
||||
model_y.fit(X, yy)
|
||||
|
||||
return model_x, model_y
|
||||
|
||||
model_lx, model_ly = train_model(left_data)
|
||||
model_rx, model_ry = train_model(right_data)
|
||||
|
||||
print("Kalibrierung abgeschlossen. Tracking startet...")
|
||||
|
||||
# Datenaufzeichnung
|
||||
records = []
|
||||
|
||||
while True:
|
||||
ret, frame = cap.read()
|
||||
frame = cv2.flip(frame, 1)
|
||||
rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
|
||||
results = face_mesh.process(rgb)
|
||||
|
||||
if results.multi_face_landmarks:
|
||||
mesh = results.multi_face_landmarks[0].landmark
|
||||
|
||||
left_center = get_iris_center(mesh, LEFT_IRIS)
|
||||
right_center = get_iris_center(mesh, RIGHT_IRIS)
|
||||
|
||||
left_input = np.array([left_center])
|
||||
right_input = np.array([right_center])
|
||||
|
||||
lx = model_lx.predict(left_input)[0]
|
||||
ly = model_ly.predict(left_input)[0]
|
||||
|
||||
rx = model_rx.predict(right_input)[0]
|
||||
ry = model_ry.predict(right_input)[0]
|
||||
|
||||
# Pixel-Koordinaten begrenzen
|
||||
lx = np.clip(lx, 0, screen_w)
|
||||
ly = np.clip(ly, 0, screen_h)
|
||||
rx = np.clip(rx, 0, screen_w)
|
||||
ry = np.clip(ry, 0, screen_h)
|
||||
|
||||
# Normierung 0–1
|
||||
lx_norm = lx / screen_w
|
||||
ly_norm = ly / screen_h
|
||||
rx_norm = rx / screen_w
|
||||
ry_norm = ry / screen_h
|
||||
|
||||
records.append([
|
||||
lx_norm, ly_norm,
|
||||
rx_norm, ry_norm
|
||||
])
|
||||
|
||||
print("L:", int(lx), int(ly), " | R:", int(rx), int(ry))
|
||||
|
||||
cv2.imshow("Tracking", frame)
|
||||
|
||||
key = cv2.waitKey(1) & 0xFF
|
||||
if key == ord('q'):
|
||||
print("q gedrückt – beende Tracking")
|
||||
break
|
||||
|
||||
cap.release()
|
||||
cv2.destroyAllWindows()
|
||||
|
||||
# CSV speichern
|
||||
df = pd.DataFrame(records, columns=[
|
||||
"EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_X",
|
||||
"EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_Y",
|
||||
"EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_X",
|
||||
"EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_Y"
|
||||
])
|
||||
|
||||
df.to_csv("gaze_data1.csv", index=False)
|
||||
|
||||
print("Daten gespeichert als gaze_data1.csv")
|
||||
@@ -0,0 +1,205 @@
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
from pathlib import Path
|
||||
from sklearn.preprocessing import MinMaxScaler
|
||||
from scipy.signal import welch
|
||||
from pygazeanalyser.detectors import fixation_detection, saccade_detection
|
||||
|
||||
|
||||
##############################################################################
|
||||
# KONFIGURATION
|
||||
##############################################################################
|
||||
|
||||
SAMPLING_RATE = 25 # Hz
|
||||
MIN_DUR_BLINKS = 2 # x * 40ms
|
||||
|
||||
|
||||
##############################################################################
|
||||
# EYE-TRACKING FUNKTIONEN
|
||||
##############################################################################
|
||||
|
||||
def clean_eye_df(df):
|
||||
"""Extrahiert nur Eye-Tracking Spalten und entfernt leere Zeilen."""
|
||||
eye_cols = [c for c in df.columns if c.startswith("EYE_")]
|
||||
|
||||
if not eye_cols:
|
||||
return pd.DataFrame()
|
||||
|
||||
df_eye = df[eye_cols].copy()
|
||||
df_eye = df_eye.replace([np.inf, -np.inf], np.nan)
|
||||
df_eye = df_eye.dropna(subset=eye_cols, how="all")
|
||||
|
||||
return df_eye.reset_index(drop=True)
|
||||
|
||||
|
||||
def extract_gaze_signal(df):
|
||||
"""Extrahiert 2D-Gaze-Positionen, maskiert ungültige Samples und interpoliert."""
|
||||
gx_L = df["EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_X"].astype(float).copy()
|
||||
gy_L = df["EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_Y"].astype(float).copy()
|
||||
gx_R = df["EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_X"].astype(float).copy()
|
||||
gy_R = df["EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_Y"].astype(float).copy()
|
||||
|
||||
val_L = (df["EYE_LEFT_PUPIL_VALIDITY"] == 1)
|
||||
val_R = (df["EYE_RIGHT_PUPIL_VALIDITY"] == 1)
|
||||
|
||||
# Inf → NaN
|
||||
for arr in [gx_L, gy_L, gx_R, gy_R]:
|
||||
arr.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
|
||||
# Ungültige maskieren
|
||||
gx_L[~val_L] = np.nan
|
||||
gy_L[~val_L] = np.nan
|
||||
gx_R[~val_R] = np.nan
|
||||
gy_R[~val_R] = np.nan
|
||||
|
||||
|
||||
# Mittelwert beider Augen
|
||||
gx = np.mean(np.column_stack([gx_L, gx_R]), axis=1)
|
||||
gy = np.mean(np.column_stack([gy_L, gy_R]), axis=1)
|
||||
|
||||
# Interpolation
|
||||
gx = pd.Series(gx).interpolate(limit=None, limit_direction="both").bfill().ffill()
|
||||
gy = pd.Series(gy).interpolate(limit=None, limit_direction="both").bfill().ffill()
|
||||
|
||||
# MinMax Skalierung
|
||||
xscaler = MinMaxScaler()
|
||||
gxscale = xscaler.fit_transform(gx.values.reshape(-1, 1))
|
||||
|
||||
yscaler = MinMaxScaler()
|
||||
gyscale = yscaler.fit_transform(gy.values.reshape(-1, 1))
|
||||
|
||||
return np.column_stack((gxscale, gyscale))
|
||||
|
||||
|
||||
def extract_pupil(df):
|
||||
"""Extrahiert Pupillengröße (beide Augen gemittelt)."""
|
||||
pl = df["EYE_LEFT_PUPIL_DIAMETER"].replace([np.inf, -np.inf], np.nan)
|
||||
pr = df["EYE_RIGHT_PUPIL_DIAMETER"].replace([np.inf, -np.inf], np.nan)
|
||||
|
||||
vl = df.get("EYE_LEFT_PUPIL_VALIDITY")
|
||||
vr = df.get("EYE_RIGHT_PUPIL_VALIDITY")
|
||||
|
||||
if vl is None or vr is None:
|
||||
validity = (~pl.isna() | ~pr.isna()).astype(int).to_numpy()
|
||||
else:
|
||||
validity = ((vl == 1) | (vr == 1)).astype(int).to_numpy()
|
||||
|
||||
p = np.mean(np.column_stack([pl, pr]), axis=1)
|
||||
p = pd.Series(p).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
|
||||
return p.to_numpy(), validity
|
||||
|
||||
|
||||
def detect_blinks(pupil_validity, min_duration=5):
|
||||
"""Erkennt Blinks: Validity=0 → Blink."""
|
||||
blinks = []
|
||||
start = None
|
||||
|
||||
for i, v in enumerate(pupil_validity):
|
||||
if v == 0 and start is None:
|
||||
start = i
|
||||
elif v == 1 and start is not None:
|
||||
if i - start >= min_duration:
|
||||
blinks.append([start, i])
|
||||
start = None
|
||||
|
||||
return blinks
|
||||
|
||||
|
||||
def compute_IPA(pupil, fs=25):
|
||||
"""Index of Pupillary Activity (Duchowski 2018)."""
|
||||
f, Pxx = welch(pupil, fs=fs, nperseg=int(fs*2))
|
||||
hf_band = (f >= 0.6) & (f <= 2.0)
|
||||
return np.sum(Pxx[hf_band])
|
||||
|
||||
|
||||
def extract_eye_features(df_eye, fs=25, min_dur_blinks=2):
|
||||
"""
|
||||
Extrahiert Eye-Tracking Features für ein einzelnes Window.
|
||||
Gibt Dictionary mit allen Eye-Features zurück.
|
||||
"""
|
||||
# Gaze
|
||||
gaze = extract_gaze_signal(df_eye)
|
||||
|
||||
# Pupille
|
||||
pupil, pupil_validity = extract_pupil(df_eye)
|
||||
|
||||
|
||||
# ----------------------------
|
||||
# FIXATIONS
|
||||
# ----------------------------
|
||||
time_ms = np.arange(len(df_eye)) * 1000.0 / fs
|
||||
|
||||
fix, efix = fixation_detection(
|
||||
x=gaze[:, 0], y=gaze[:, 1], time=time_ms,
|
||||
missing=0.0, maxdist=0.003, mindur=10
|
||||
)
|
||||
|
||||
fixation_durations = [f[2] for f in efix if np.isfinite(f[2]) and f[2] > 0]
|
||||
|
||||
# Kategorien
|
||||
F_short = sum(66 <= d <= 150 for d in fixation_durations)
|
||||
F_medium = sum(300 <= d <= 500 for d in fixation_durations)
|
||||
F_long = sum(d >= 1000 for d in fixation_durations)
|
||||
F_hundred = sum(d > 100 for d in fixation_durations)
|
||||
|
||||
# ----------------------------
|
||||
# SACCADES
|
||||
# ----------------------------
|
||||
sac, esac = saccade_detection(
|
||||
x=gaze[:, 0], y=gaze[:, 1], time=time_ms,
|
||||
missing=0, minlen=12, maxvel=0.2, maxacc=1
|
||||
)
|
||||
|
||||
sac_durations = [s[2] for s in esac]
|
||||
sac_amplitudes = [((s[5]-s[3])**2 + (s[6]-s[4])**2)**0.5 for s in esac]
|
||||
|
||||
# ----------------------------
|
||||
# BLINKS
|
||||
# ----------------------------
|
||||
blinks = detect_blinks(pupil_validity, min_duration=min_dur_blinks)
|
||||
blink_durations = [(b[1] - b[0]) / fs for b in blinks]
|
||||
|
||||
# ----------------------------
|
||||
# PUPIL
|
||||
# ----------------------------
|
||||
if np.all(np.isnan(pupil)):
|
||||
mean_pupil = np.nan
|
||||
ipa = np.nan
|
||||
else:
|
||||
mean_pupil = np.nanmean(pupil)
|
||||
ipa = compute_IPA(pupil, fs=fs)
|
||||
|
||||
# Feature Dictionary
|
||||
return {
|
||||
"Fix_count_short_66_150": F_short,
|
||||
"Fix_count_medium_300_500": F_medium,
|
||||
"Fix_count_long_gt_1000": F_long,
|
||||
"Fix_count_100": F_hundred,
|
||||
"Fix_mean_duration": np.mean(fixation_durations) if fixation_durations else 0,
|
||||
"Fix_median_duration": np.median(fixation_durations) if fixation_durations else 0,
|
||||
"Sac_count": len(sac),
|
||||
"Sac_mean_amp": np.mean(sac_amplitudes) if sac_amplitudes else 0,
|
||||
"Sac_mean_dur": np.mean(sac_durations) if sac_durations else 0,
|
||||
"Sac_median_dur": np.median(sac_durations) if sac_durations else 0,
|
||||
"Blink_count": len(blinks),
|
||||
"Blink_mean_dur": np.mean(blink_durations) if blink_durations else 0,
|
||||
"Blink_median_dur": np.median(blink_durations) if blink_durations else 0,
|
||||
"Pupil_mean": mean_pupil,
|
||||
"Pupil_IPA": ipa
|
||||
}
|
||||
|
||||
def compute_features_from_parquet(parquet_path):
|
||||
df = pd.read_parquet(parquet_path)
|
||||
df_eye = clean_eye_df(df)
|
||||
|
||||
if df_eye.empty:
|
||||
return None
|
||||
|
||||
features = extract_eye_features(
|
||||
df_eye,
|
||||
fs=SAMPLING_RATE,
|
||||
min_dur_blinks=MIN_DUR_BLINKS
|
||||
)
|
||||
|
||||
return features
|
||||
@@ -1,91 +0,0 @@
|
||||
import os
|
||||
import pandas as pd
|
||||
from pathlib import Path
|
||||
|
||||
print(os.getcwd())
|
||||
num_files = 2 # number of files to process (min: 1, max: 30)
|
||||
|
||||
print("connection aufgebaut")
|
||||
|
||||
data_dir = Path("/home/jovyan/Fahrsimulator_MSY2526_AI/EDA")
|
||||
# os.chdir(data_dir)
|
||||
# Get all .h5 files and sort them
|
||||
matching_files = sorted(data_dir.glob("*.h5"))
|
||||
|
||||
# Chunk size for reading (adjust based on your RAM - 100k rows is ~50-100MB depending on columns)
|
||||
CHUNK_SIZE = 100_000
|
||||
|
||||
for i, file_path in enumerate(matching_files):
|
||||
print(f"Subject {i} gestartet")
|
||||
print(f"{file_path} geoeffnet")
|
||||
|
||||
# Step 1: Get total number of rows and column names
|
||||
with pd.HDFStore(file_path, mode="r") as store:
|
||||
cols = store.select("SIGNALS", start=0, stop=1).columns
|
||||
nrows = store.get_storer("SIGNALS").nrows
|
||||
print(f"Total columns: {len(cols)}, Total rows: {nrows}")
|
||||
|
||||
# Step 2: Filter columns that start with "FACE_AU"
|
||||
eye_cols = [c for c in cols if c.startswith("EYE_")]
|
||||
print(f"eye-tracking columns found: {eye_cols}")
|
||||
|
||||
if len(eye_cols) == 0:
|
||||
print(f"keine eye-tracking-Signale in Subject {i}")
|
||||
continue
|
||||
|
||||
# Columns to read
|
||||
columns_to_read = ["STUDY", "LEVEL", "PHASE"] + eye_cols
|
||||
|
||||
# Step 3: Process file in chunks
|
||||
chunks_to_save = []
|
||||
|
||||
for start_row in range(0, nrows, CHUNK_SIZE):
|
||||
stop_row = min(start_row + CHUNK_SIZE, nrows)
|
||||
print(f"Processing rows {start_row} to {stop_row} ({stop_row/nrows*100:.1f}%)")
|
||||
|
||||
# Read chunk
|
||||
df_chunk = pd.read_hdf(
|
||||
file_path,
|
||||
key="SIGNALS",
|
||||
columns=columns_to_read,
|
||||
start=start_row,
|
||||
stop=stop_row
|
||||
)
|
||||
|
||||
# Add metadata columns
|
||||
df_chunk["subjectID"] = i
|
||||
df_chunk["rowID"] = range(start_row, stop_row)
|
||||
|
||||
# Clean data
|
||||
df_chunk = df_chunk[df_chunk["LEVEL"] != 0]
|
||||
df_chunk = df_chunk.dropna()
|
||||
|
||||
# Only keep non-empty chunks
|
||||
if len(df_chunk) > 0:
|
||||
chunks_to_save.append(df_chunk)
|
||||
|
||||
# Free memory
|
||||
del df_chunk
|
||||
|
||||
print("load and cleaning done")
|
||||
|
||||
# Step 4: Combine all chunks and save
|
||||
if chunks_to_save:
|
||||
df_final = pd.concat(chunks_to_save, ignore_index=True)
|
||||
print(f"Final dataframe shape: {df_final.shape}")
|
||||
|
||||
# Save to parquet
|
||||
base_dir = Path(r"/home/jovyan/data-paulusjafahrsimulator-gpu/new_ET_Parquet_files")
|
||||
os.makedirs(base_dir, exist_ok=True)
|
||||
|
||||
out_name = base_dir / f"ET_signals_extracted_{i:04d}.parquet"
|
||||
df_final.to_parquet(out_name, index=False)
|
||||
print(f"Saved to {out_name}")
|
||||
|
||||
# Free memory
|
||||
del df_final
|
||||
del chunks_to_save
|
||||
else:
|
||||
print(f"No valid data found for Subject {i}")
|
||||
|
||||
print("All files processed!")
|
||||
@@ -1,91 +0,0 @@
|
||||
import os
|
||||
import pandas as pd
|
||||
from pathlib import Path
|
||||
|
||||
print(os.getcwd())
|
||||
num_files = 2 # number of files to process (min: 1, max: 30)
|
||||
|
||||
print("connection aufgebaut")
|
||||
|
||||
data_dir = Path(r"C:\Users\x\repo\UXKI\Fahrsimulator_MSY2526_AI\newTmp")
|
||||
|
||||
# Get all .h5 files and sort them
|
||||
matching_files = sorted(data_dir.glob("*.h5"))
|
||||
|
||||
# Chunk size for reading (adjust based on your RAM - 100k rows is ~50-100MB depending on columns)
|
||||
CHUNK_SIZE = 100_000
|
||||
|
||||
for i, file_path in enumerate(matching_files):
|
||||
print(f"Subject {i} gestartet")
|
||||
print(f"{file_path} geoeffnet")
|
||||
|
||||
# Step 1: Get total number of rows and column names
|
||||
with pd.HDFStore(file_path, mode="r") as store:
|
||||
cols = store.select("SIGNALS", start=0, stop=1).columns
|
||||
nrows = store.get_storer("SIGNALS").nrows
|
||||
print(f"Total columns: {len(cols)}, Total rows: {nrows}")
|
||||
|
||||
# Step 2: Filter columns that start with "FACE_AU"
|
||||
eye_cols = [c for c in cols if c.startswith("FACE_AU")]
|
||||
print(f"FACE_AU columns found: {eye_cols}")
|
||||
|
||||
if len(eye_cols) == 0:
|
||||
print(f"keine FACE_AU-Signale in Subject {i}")
|
||||
continue
|
||||
|
||||
# Columns to read
|
||||
columns_to_read = ["STUDY", "LEVEL", "PHASE"] + eye_cols
|
||||
|
||||
# Step 3: Process file in chunks
|
||||
chunks_to_save = []
|
||||
|
||||
for start_row in range(0, nrows, CHUNK_SIZE):
|
||||
stop_row = min(start_row + CHUNK_SIZE, nrows)
|
||||
print(f"Processing rows {start_row} to {stop_row} ({stop_row/nrows*100:.1f}%)")
|
||||
|
||||
# Read chunk
|
||||
df_chunk = pd.read_hdf(
|
||||
file_path,
|
||||
key="SIGNALS",
|
||||
columns=columns_to_read,
|
||||
start=start_row,
|
||||
stop=stop_row
|
||||
)
|
||||
|
||||
# Add metadata columns
|
||||
df_chunk["subjectID"] = i
|
||||
df_chunk["rowID"] = range(start_row, stop_row)
|
||||
|
||||
# Clean data
|
||||
df_chunk = df_chunk[df_chunk["LEVEL"] != 0]
|
||||
df_chunk = df_chunk.dropna()
|
||||
|
||||
# Only keep non-empty chunks
|
||||
if len(df_chunk) > 0:
|
||||
chunks_to_save.append(df_chunk)
|
||||
|
||||
# Free memory
|
||||
del df_chunk
|
||||
|
||||
print("load and cleaning done")
|
||||
|
||||
# Step 4: Combine all chunks and save
|
||||
if chunks_to_save:
|
||||
df_final = pd.concat(chunks_to_save, ignore_index=True)
|
||||
print(f"Final dataframe shape: {df_final.shape}")
|
||||
|
||||
# Save to parquet
|
||||
base_dir = Path(r"C:\new_AU_parquet_files")
|
||||
os.makedirs(base_dir, exist_ok=True)
|
||||
|
||||
out_name = base_dir / f"cleaned_{i:04d}.parquet"
|
||||
df_final.to_parquet(out_name, index=False)
|
||||
print(f"Saved to {out_name}")
|
||||
|
||||
# Free memory
|
||||
del df_final
|
||||
del chunks_to_save
|
||||
else:
|
||||
print(f"No valid data found for Subject {i}")
|
||||
|
||||
print("All files processed!")
|
||||
@@ -4,27 +4,28 @@ import pandas as pd
|
||||
from pathlib import Path
|
||||
from sklearn.preprocessing import MinMaxScaler
|
||||
from scipy.signal import welch
|
||||
from pygazeanalyser.detectors import fixation_detection, saccade_detection
|
||||
from pygazeanalyser.detectors import fixation_detection, saccade_detection # not installed by default
|
||||
|
||||
|
||||
##############################################################################
|
||||
# KONFIGURATION
|
||||
# CONFIGURATION
|
||||
##############################################################################
|
||||
INPUT_DIR = Path(r"/home/jovyan/data-paulusjafahrsimulator-gpu/both_mod_parquet_files")
|
||||
OUTPUT_FILE = Path(r"/home/jovyan/data-paulusjafahrsimulator-gpu/new_datasets/combined_dataset_25hz.parquet")
|
||||
|
||||
WINDOW_SIZE_SAMPLES = 1250 # 50s bei 25Hz
|
||||
STEP_SIZE_SAMPLES = 125 # 5s bei 25Hz
|
||||
INPUT_DIR = Path(r"") # directory that stores the parquet files (one file per subject)
|
||||
OUTPUT_FILE = Path(r"") # path for resulting dataset
|
||||
WINDOW_SIZE_SAMPLES = 25*50 # 50s at 25Hz
|
||||
STEP_SIZE_SAMPLES = 125 # 5s at 25Hz
|
||||
SAMPLING_RATE = 25 # Hz
|
||||
MIN_DUR_BLINKS = 2 # x * 40ms
|
||||
|
||||
|
||||
##############################################################################
|
||||
# EYE-TRACKING FUNKTIONEN
|
||||
# EYE-TRACKING FUNCTIONS
|
||||
##############################################################################
|
||||
|
||||
def clean_eye_df(df):
|
||||
"""Extrahiert nur Eye-Tracking Spalten und entfernt leere Zeilen."""
|
||||
"""Extracts Eye-Tracking columns only and removes empty rows."""
|
||||
eye_cols = [c for c in df.columns if c.startswith("EYE_")]
|
||||
|
||||
if not eye_cols:
|
||||
return pd.DataFrame()
|
||||
|
||||
@@ -36,7 +37,7 @@ def clean_eye_df(df):
|
||||
|
||||
|
||||
def extract_gaze_signal(df):
|
||||
"""Extrahiert 2D-Gaze-Positionen, maskiert ungültige Samples und interpoliert."""
|
||||
"""Extracts 2D gaze positions, masks invalid samples, and interpolates."""
|
||||
gx_L = df["EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_X"].astype(float).copy()
|
||||
gy_L = df["EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_Y"].astype(float).copy()
|
||||
gx_R = df["EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_X"].astype(float).copy()
|
||||
@@ -49,21 +50,22 @@ def extract_gaze_signal(df):
|
||||
for arr in [gx_L, gy_L, gx_R, gy_R]:
|
||||
arr.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
|
||||
# Ungültige maskieren
|
||||
# Mask invalids
|
||||
gx_L[~val_L] = np.nan
|
||||
gy_L[~val_L] = np.nan
|
||||
gx_R[~val_R] = np.nan
|
||||
gy_R[~val_R] = np.nan
|
||||
|
||||
# Mittelwert beider Augen
|
||||
|
||||
# Mean of both eyes
|
||||
gx = np.mean(np.column_stack([gx_L, gx_R]), axis=1)
|
||||
gy = np.mean(np.column_stack([gy_L, gy_R]), axis=1)
|
||||
|
||||
# Interpolation
|
||||
gx = pd.Series(gx).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
gy = pd.Series(gy).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
gx = pd.Series(gx).interpolate(limit=None, limit_direction="both").bfill().ffill()
|
||||
gy = pd.Series(gy).interpolate(limit=None, limit_direction="both").bfill().ffill()
|
||||
|
||||
# MinMax Skalierung
|
||||
# MinMax scaling
|
||||
xscaler = MinMaxScaler()
|
||||
gxscale = xscaler.fit_transform(gx.values.reshape(-1, 1))
|
||||
|
||||
@@ -74,7 +76,7 @@ def extract_gaze_signal(df):
|
||||
|
||||
|
||||
def extract_pupil(df):
|
||||
"""Extrahiert Pupillengröße (beide Augen gemittelt)."""
|
||||
"""Extract pupil size (average of both eyes)."""
|
||||
pl = df["EYE_LEFT_PUPIL_DIAMETER"].replace([np.inf, -np.inf], np.nan)
|
||||
pr = df["EYE_RIGHT_PUPIL_DIAMETER"].replace([np.inf, -np.inf], np.nan)
|
||||
|
||||
@@ -93,7 +95,7 @@ def extract_pupil(df):
|
||||
|
||||
|
||||
def detect_blinks(pupil_validity, min_duration=5):
|
||||
"""Erkennt Blinks: Validity=0 → Blink."""
|
||||
"""Detect blinks: Validity=0 → Blink."""
|
||||
blinks = []
|
||||
start = None
|
||||
|
||||
@@ -115,15 +117,15 @@ def compute_IPA(pupil, fs=25):
|
||||
return np.sum(Pxx[hf_band])
|
||||
|
||||
|
||||
def extract_eye_features_window(df_eye_window, fs=25):
|
||||
def extract_eye_features_window(df_eye_window, fs=25, min_dur_blinks=2):
|
||||
"""
|
||||
Extrahiert Eye-Tracking Features für ein einzelnes Window.
|
||||
Gibt Dictionary mit allen Eye-Features zurück.
|
||||
Extracts eye tracking features for a single window.
|
||||
Returns a dictionary containing all eye features.
|
||||
"""
|
||||
# Gaze
|
||||
gaze = extract_gaze_signal(df_eye_window)
|
||||
|
||||
# Pupille
|
||||
# Pupil
|
||||
pupil, pupil_validity = extract_pupil(df_eye_window)
|
||||
|
||||
window_size = len(df_eye_window)
|
||||
@@ -140,7 +142,6 @@ def extract_eye_features_window(df_eye_window, fs=25):
|
||||
|
||||
fixation_durations = [f[2] for f in efix if np.isfinite(f[2]) and f[2] > 0]
|
||||
|
||||
# Kategorien
|
||||
F_short = sum(66 <= d <= 150 for d in fixation_durations)
|
||||
F_medium = sum(300 <= d <= 500 for d in fixation_durations)
|
||||
F_long = sum(d >= 1000 for d in fixation_durations)
|
||||
@@ -160,7 +161,7 @@ def extract_eye_features_window(df_eye_window, fs=25):
|
||||
# ----------------------------
|
||||
# BLINKS
|
||||
# ----------------------------
|
||||
blinks = detect_blinks(pupil_validity)
|
||||
blinks = detect_blinks(pupil_validity, min_duration=min_dur_blinks)
|
||||
blink_durations = [(b[1] - b[0]) / fs for b in blinks]
|
||||
|
||||
# ----------------------------
|
||||
@@ -194,27 +195,27 @@ def extract_eye_features_window(df_eye_window, fs=25):
|
||||
|
||||
|
||||
##############################################################################
|
||||
# KOMBINIERTE FEATURE-EXTRAKTION
|
||||
# Combined feature extraction
|
||||
##############################################################################
|
||||
|
||||
def process_combined_features(input_dir, output_file, window_size, step_size, fs=25):
|
||||
def process_combined_features(input_dir, output_file, window_size, step_size, fs=25,min_duration_blinks=2):
|
||||
"""
|
||||
Verarbeitet Parquet-Dateien mit FACE_AU und EYE Spalten.
|
||||
Extrahiert beide Feature-Sets und kombiniert sie.
|
||||
Processes Parquet files with FACE_AU and EYE columns.
|
||||
Extracts both feature sets and combines them.
|
||||
"""
|
||||
input_path = Path(input_dir)
|
||||
parquet_files = sorted(input_path.glob("*.parquet"))
|
||||
|
||||
if not parquet_files:
|
||||
print(f"FEHLER: Keine Parquet-Dateien in {input_dir} gefunden!")
|
||||
print(f"Error: No parquet-files found in {input_dir}!")
|
||||
return None
|
||||
|
||||
print(f"\n{'='*70}")
|
||||
print(f"KOMBINIERTE FEATURE-EXTRAKTION")
|
||||
print(f"Combined feature-extraction")
|
||||
print(f"{'='*70}")
|
||||
print(f"Dateien: {len(parquet_files)}")
|
||||
print(f"Window: {window_size} Samples ({window_size/fs:.1f}s bei {fs}Hz)")
|
||||
print(f"Step: {step_size} Samples ({step_size/fs:.1f}s bei {fs}Hz)")
|
||||
print(f"Files: {len(parquet_files)}")
|
||||
print(f"Window: {window_size} Samples ({window_size/fs:.1f}s at {fs}Hz)")
|
||||
print(f"Step: {step_size} Samples ({step_size/fs:.1f}s at {fs}Hz)")
|
||||
print(f"{'='*70}\n")
|
||||
|
||||
all_windows = []
|
||||
@@ -224,23 +225,22 @@ def process_combined_features(input_dir, output_file, window_size, step_size, fs
|
||||
|
||||
try:
|
||||
df = pd.read_parquet(parquet_file)
|
||||
print(f" Einträge: {len(df)}")
|
||||
print(f" Entries: {len(df)}")
|
||||
|
||||
# Identifiziere Spalten
|
||||
au_columns = [col for col in df.columns if col.startswith('FACE_AU')]
|
||||
eye_columns = [col for col in df.columns if col.startswith('EYE_')]
|
||||
|
||||
print(f" AU-Spalten: {len(au_columns)}")
|
||||
print(f" Eye-Spalten: {len(eye_columns)}")
|
||||
print(f" AU-columns: {len(au_columns)}")
|
||||
print(f" Eye-columns: {len(eye_columns)}")
|
||||
|
||||
has_au = len(au_columns) > 0
|
||||
has_eye = len(eye_columns) > 0
|
||||
|
||||
if not has_au and not has_eye:
|
||||
print(f" WARNUNG: Keine AU oder Eye Spalten gefunden!")
|
||||
print(f" Warning: No AU or eye tracking columns found!")
|
||||
continue
|
||||
|
||||
# Gruppiere nach STUDY, LEVEL, PHASE
|
||||
# Group by STUDY, LEVEL, PHASE
|
||||
group_cols = [col for col in ['STUDY', 'LEVEL', 'PHASE'] if col in df.columns]
|
||||
|
||||
if group_cols:
|
||||
@@ -254,7 +254,7 @@ def process_combined_features(input_dir, output_file, window_size, step_size, fs
|
||||
|
||||
group_df = group_df.reset_index(drop=True)
|
||||
|
||||
# Berechne Anzahl Windows
|
||||
# calculate number of windows
|
||||
num_windows = (len(group_df) - window_size) // step_size + 1
|
||||
|
||||
if num_windows <= 0:
|
||||
@@ -268,7 +268,7 @@ def process_combined_features(input_dir, output_file, window_size, step_size, fs
|
||||
|
||||
window_df = group_df.iloc[start_idx:end_idx]
|
||||
|
||||
# Basis-Metadaten
|
||||
# basic metadata
|
||||
result = {
|
||||
'subjectID': window_df['subjectID'].iloc[0],
|
||||
'start_time': window_df['rowID'].iloc[0],
|
||||
@@ -277,19 +277,22 @@ def process_combined_features(input_dir, output_file, window_size, step_size, fs
|
||||
'PHASE': window_df['PHASE'].iloc[0] if 'PHASE' in window_df.columns else np.nan
|
||||
}
|
||||
|
||||
# FACE AU Features
|
||||
# FACE AU features
|
||||
if has_au:
|
||||
for au_col in au_columns:
|
||||
result[f'{au_col}_mean'] = window_df[au_col].mean()
|
||||
|
||||
# Eye-Tracking Features
|
||||
# Eye-tracking features
|
||||
if has_eye:
|
||||
try:
|
||||
eye_features = extract_eye_features_window(window_df[eye_columns], fs=fs)
|
||||
# clean dataframe from all nan rows
|
||||
window_df= clean_eye_df(window_df)
|
||||
|
||||
eye_features = extract_eye_features_window(window_df[eye_columns], fs=fs,min_dur_blinks=min_duration_blinks)
|
||||
result.update(eye_features)
|
||||
except Exception as e:
|
||||
print(f" WARNUNG: Eye-Features fehlgeschlagen: {str(e)}")
|
||||
# Füge NaN-Werte für Eye-Features hinzu
|
||||
# Add NaN-values for eye-features
|
||||
result.update({
|
||||
"Fix_count_short_66_150": np.nan,
|
||||
"Fix_count_medium_300_500": np.nan,
|
||||
@@ -318,7 +321,7 @@ def process_combined_features(input_dir, output_file, window_size, step_size, fs
|
||||
traceback.print_exc()
|
||||
continue
|
||||
|
||||
# Kombiniere alle Windows
|
||||
# Combine all windows
|
||||
if not all_windows:
|
||||
print("\nKEINE FEATURES EXTRAHIERT!")
|
||||
return None
|
||||
@@ -333,7 +336,7 @@ def process_combined_features(input_dir, output_file, window_size, step_size, fs
|
||||
print(f"Spalten: {len(result_df.columns)}")
|
||||
print(f"Subjects: {result_df['subjectID'].nunique()}")
|
||||
|
||||
# Speichern
|
||||
# Save
|
||||
output_path = Path(output_file)
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
result_df.to_parquet(output_file, index=False)
|
||||
@@ -350,7 +353,7 @@ def process_combined_features(input_dir, output_file, window_size, step_size, fs
|
||||
|
||||
def main():
|
||||
print("\n" + "="*70)
|
||||
print("KOMBINIERTE FEATURE-EXTRAKTION (AU + EYE)")
|
||||
print("Combined extraction (AU + EYE)")
|
||||
print("="*70)
|
||||
|
||||
result = process_combined_features(
|
||||
@@ -358,23 +361,21 @@ def main():
|
||||
output_file=OUTPUT_FILE,
|
||||
window_size=WINDOW_SIZE_SAMPLES,
|
||||
step_size=STEP_SIZE_SAMPLES,
|
||||
fs=SAMPLING_RATE
|
||||
fs=SAMPLING_RATE,
|
||||
min_duration_blinks=MIN_DUR_BLINKS
|
||||
)
|
||||
|
||||
if result is not None:
|
||||
print("\nErste 5 Zeilen:")
|
||||
print("\First 5 rows:")
|
||||
print(result.head())
|
||||
|
||||
print("\nSpalten-Übersicht:")
|
||||
print(result.columns.tolist())
|
||||
|
||||
print("\nDatentypen:")
|
||||
print("\nColumns overview:")
|
||||
print(result.dtypes)
|
||||
|
||||
print("\nStatistik:")
|
||||
print("\Statistics:")
|
||||
print(result.describe())
|
||||
|
||||
print("\n✓ FERTIG!\n")
|
||||
print("\nDone!\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,113 +0,0 @@
|
||||
import pandas as pd
|
||||
import numpy as np
|
||||
from pathlib import Path
|
||||
|
||||
def process_parquet_files(input_dir, output_file, window_size=1250, step_size=125):
|
||||
"""
|
||||
Verarbeitet Parquet-Dateien mit Sliding Window Aggregation.
|
||||
|
||||
Parameters:
|
||||
-----------
|
||||
input_dir : str
|
||||
Verzeichnis mit Parquet-Dateien
|
||||
output_file : str
|
||||
Pfad für die Ausgabe-Parquet-Datei
|
||||
window_size : int
|
||||
Größe des Sliding Windows (default: 3000)
|
||||
step_size : int
|
||||
Schrittweite in Einträgen (default: 250 = 10 Sekunden bei 25 Hz)
|
||||
"""
|
||||
|
||||
input_path = Path(input_dir)
|
||||
parquet_files = sorted(input_path.glob("*.parquet"))
|
||||
|
||||
if not parquet_files:
|
||||
print(f"Keine Parquet-Dateien in {input_dir} gefunden!")
|
||||
return
|
||||
|
||||
print(f"Gefundene Dateien: {len(parquet_files)}")
|
||||
|
||||
all_windows = []
|
||||
|
||||
for file_idx, parquet_file in enumerate(parquet_files):
|
||||
print(f"\nVerarbeite Datei {file_idx + 1}/{len(parquet_files)}: {parquet_file.name}")
|
||||
|
||||
# Lade Parquet-Datei
|
||||
df = pd.read_parquet(parquet_file)
|
||||
print(f" Einträge: {len(df)}")
|
||||
|
||||
# Identifiziere AU-Spalten
|
||||
au_columns = [col for col in df.columns if col.startswith('FACE_AU')]
|
||||
print(f" AU-Spalten: {len(au_columns)}")
|
||||
|
||||
# Gruppiere nach STUDY, LEVEL, PHASE (um Übergänge zu vermeiden)
|
||||
for (study_val, level_val, phase_val), level_df in df.groupby(['STUDY', 'LEVEL', 'PHASE'], sort=False):
|
||||
print(f" STUDY {study_val}, LEVEL {level_val}, PHASE {phase_val}: {len(level_df)} Einträge")
|
||||
|
||||
# Reset index für korrekte Position-Berechnung
|
||||
level_df = level_df.reset_index(drop=True)
|
||||
|
||||
# Sliding Window über dieses Level
|
||||
num_windows = (len(level_df) - window_size) // step_size + 1
|
||||
|
||||
if num_windows <= 0:
|
||||
print(f" Zu wenige Einträge für Window (benötigt {window_size})")
|
||||
continue
|
||||
|
||||
for i in range(num_windows):
|
||||
start_idx = i * step_size
|
||||
end_idx = start_idx + window_size
|
||||
|
||||
window_df = level_df.iloc[start_idx:end_idx]
|
||||
|
||||
# Erstelle aggregiertes Ergebnis
|
||||
result = {
|
||||
'subjectID': window_df['subjectID'].iloc[0],
|
||||
'start_time': window_df['rowID'].iloc[0], # rowID als start_time
|
||||
'STUDY': window_df['STUDY'].iloc[0],
|
||||
'LEVEL': window_df['LEVEL'].iloc[0],
|
||||
'PHASE': window_df['PHASE'].iloc[0]
|
||||
}
|
||||
|
||||
# Summiere alle AU-Spalten
|
||||
for au_col in au_columns:
|
||||
# result[f'{au_col}_sum'] = window_df[au_col].sum()
|
||||
result[f'{au_col}_mean'] = window_df[au_col].mean()
|
||||
|
||||
all_windows.append(result)
|
||||
|
||||
print(f" Windows erstellt: {num_windows}")
|
||||
|
||||
# Erstelle finalen DataFrame
|
||||
result_df = pd.DataFrame(all_windows)
|
||||
|
||||
print(f"\n{'='*60}")
|
||||
print(f"Gesamt Windows erstellt: {len(result_df)}")
|
||||
print(f"Spalten: {list(result_df.columns)}")
|
||||
|
||||
# Speichere Ergebnis
|
||||
result_df.to_parquet(output_file, index=False)
|
||||
print(f"\nErgebnis gespeichert in: {output_file}")
|
||||
|
||||
return result_df
|
||||
|
||||
|
||||
# Beispiel-Verwendung
|
||||
if __name__ == "__main__":
|
||||
# Anpassen an deine Pfade
|
||||
input_directory = Path(r"/home/jovyan/data-paulusjafahrsimulator-gpu/new_AU_parquet_files")
|
||||
output_file = Path(r"/home/jovyan/data-paulusjafahrsimulator-gpu/new_AU_dataset_mean/AU_dataset_mean.parquet")
|
||||
|
||||
|
||||
|
||||
result = process_parquet_files(
|
||||
input_dir=input_directory,
|
||||
output_file=output_file,
|
||||
window_size=1250,
|
||||
step_size=125
|
||||
)
|
||||
|
||||
# Zeige erste Zeilen
|
||||
if result is not None:
|
||||
print("\nErste 5 Zeilen des Ergebnisses:")
|
||||
print(result.head())
|
||||
@@ -1,56 +0,0 @@
|
||||
from pathlib import Path
|
||||
import pandas as pd
|
||||
|
||||
|
||||
def main():
|
||||
"""
|
||||
USER CONFIGURATION
|
||||
------------------
|
||||
Specify input files and output directory here.
|
||||
"""
|
||||
|
||||
# Input parquet files (single-modality datasets)
|
||||
file_modality_1 = Path("/home/jovyan/data-paulusjafahrsimulator-gpu/new_datasets/AU_dataset_mean.parquet")
|
||||
file_modality_2 = Path("/home/jovyan/data-paulusjafahrsimulator-gpu/new_datasets/new_eye_dataset.parquet")
|
||||
|
||||
# Output directory and file name
|
||||
output_dir = Path("/home/jovyan/data-paulusjafahrsimulator-gpu/new_datasets/")
|
||||
output_file = output_dir / "merged_dataset.parquet"
|
||||
|
||||
# Column names (adjust only if your schema differs)
|
||||
subject_col = "subjectID"
|
||||
time_col = "start_time"
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Load datasets
|
||||
# ------------------------------------------------------------------
|
||||
df1 = pd.read_parquet(file_modality_1)
|
||||
df2 = pd.read_parquet(file_modality_2)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Keep only subjects that appear in BOTH datasets
|
||||
# ------------------------------------------------------------------
|
||||
common_subjects = set(df1[subject_col]).intersection(df2[subject_col])
|
||||
|
||||
df1 = df1[df1[subject_col].isin(common_subjects)]
|
||||
df2 = df2[df2[subject_col].isin(common_subjects)]
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Inner join on subject ID AND start_time
|
||||
# ------------------------------------------------------------------
|
||||
merged_df = pd.merge(
|
||||
df1,
|
||||
df2,
|
||||
on=[subject_col, time_col],
|
||||
how="inner",
|
||||
)
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Save merged dataset
|
||||
# ------------------------------------------------------------------
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
merged_df.to_parquet(output_file, index=False)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+5
-11
@@ -1,6 +1,5 @@
|
||||
# pip install pyocclient
|
||||
import yaml
|
||||
import owncloud
|
||||
import owncloud # pip install pyocclient
|
||||
import pandas as pd
|
||||
import h5py
|
||||
import os
|
||||
@@ -26,7 +25,7 @@ for i in range(num_files):
|
||||
|
||||
# Download file from ownCloud
|
||||
oc.get_file(file_name, local_tmp)
|
||||
print(f"{file_name} geoeffnet")
|
||||
print(f"Opened: {file_name}")
|
||||
# Load into memory and extract needed columns
|
||||
# with h5py.File(local_tmp, "r") as f:
|
||||
# # Adjust this path depending on actual dataset layout inside .h5py file
|
||||
@@ -35,14 +34,9 @@ for i in range(num_files):
|
||||
with pd.HDFStore(local_tmp, mode="r") as store:
|
||||
cols = store.select("SIGNALS", start=0, stop=1).columns # get column names
|
||||
|
||||
# Step 2: Filter columns that start with "AU"
|
||||
au_cols = [c for c in cols if c.startswith("AU")]
|
||||
print(au_cols)
|
||||
if len(au_cols)==0:
|
||||
print(f"keine AU Signale in Subject {i}")
|
||||
continue
|
||||
|
||||
# Step 3: Read only those columns (plus any others you want)
|
||||
df = pd.read_hdf(local_tmp, key="SIGNALS", columns=["STUDY", "LEVEL", "PHASE"] + au_cols)
|
||||
df = pd.read_hdf(local_tmp, key="SIGNALS", columns=["STUDY", "LEVEL", "PHASE"] + cols)
|
||||
|
||||
|
||||
print("load done")
|
||||
@@ -63,7 +57,7 @@ for i in range(num_files):
|
||||
|
||||
|
||||
# Save to parquet
|
||||
os.makedirs("ParquetFiles", exist_ok=True)
|
||||
os.makedirs("ParquetFiles", exist_ok=True) # TODO: change for custom directory
|
||||
out_name = f"ParquetFiles/cleaned_{i:04d}.parquet"
|
||||
df.to_parquet(out_name, index=False)
|
||||
|
||||
@@ -1,323 +0,0 @@
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import h5py
|
||||
import yaml
|
||||
import os
|
||||
from sklearn.preprocessing import MinMaxScaler
|
||||
from scipy.signal import welch
|
||||
from pygazeanalyser.detectors import fixation_detection, saccade_detection
|
||||
|
||||
|
||||
##############################################################################
|
||||
# 1. HELFERFUNKTIONEN
|
||||
##############################################################################
|
||||
def clean_eye_df(df):
|
||||
"""
|
||||
Entfernt alle Zeilen, die keine echten Eyetracking-Daten enthalten.
|
||||
Löst das Problem, dass das Haupt-DataFrame NaN-Zeilen für andere Sensoren enthält.
|
||||
"""
|
||||
eye_cols = [c for c in df.columns if ("LEFT_" in c or "RIGHT_" in c)]
|
||||
df_eye = df[eye_cols]
|
||||
|
||||
# INF → NaN
|
||||
df_eye = df_eye.replace([np.inf, -np.inf], np.nan)
|
||||
|
||||
# Nur Zeilen behalten, wo es echte Eyetracking-Daten gibt
|
||||
df_eye = df_eye.dropna(subset=eye_cols, how="all")
|
||||
|
||||
print("Eyetracking-Zeilen vorher:", len(df))
|
||||
print("Eyetracking-Zeilen nachher:", len(df_eye))
|
||||
|
||||
#Index zurücksetzen
|
||||
return df_eye.reset_index(drop=True)
|
||||
|
||||
|
||||
def extract_gaze_signal(df):
|
||||
"""
|
||||
Extrahiert 2D-Gaze-Positionen auf dem Display,
|
||||
maskiert ungültige Samples und interpoliert Lücken.
|
||||
"""
|
||||
|
||||
print("→ extract_gaze_signal(): Eingabegröße:", df.shape)
|
||||
|
||||
# Gaze-Spalten
|
||||
gx_L = df["LEFT_GAZE_POINT_ON_DISPLAY_AREA_X"].astype(float).copy()
|
||||
gy_L = df["LEFT_GAZE_POINT_ON_DISPLAY_AREA_Y"].astype(float).copy()
|
||||
gx_R = df["RIGHT_GAZE_POINT_ON_DISPLAY_AREA_X"].astype(float).copy()
|
||||
gy_R = df["RIGHT_GAZE_POINT_ON_DISPLAY_AREA_Y"].astype(float).copy()
|
||||
|
||||
|
||||
# Validity-Spalten (1 = gültig)
|
||||
val_L = (df["LEFT_GAZE_POINT_VALIDITY"] == 1)
|
||||
val_R = (df["RIGHT_GAZE_POINT_VALIDITY"] == 1)
|
||||
|
||||
# Inf ersetzen mit NaN (kommt bei Tobii bei Blinks vor)
|
||||
gx_L.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
gy_L.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
gx_R.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
gy_R.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
|
||||
# Ungültige Werte maskieren
|
||||
gx_L[~val_L] = np.nan
|
||||
gy_L[~val_L] = np.nan
|
||||
gx_R[~val_R] = np.nan
|
||||
gy_R[~val_R] = np.nan
|
||||
|
||||
# Mittelwert der beiden Augen pro Sample (nanmean ist robust)
|
||||
gx = np.mean(np.column_stack([gx_L, gx_R]), axis=1)
|
||||
gy = np.mean(np.column_stack([gy_L, gy_R]), axis=1)
|
||||
|
||||
# Interpolation (wichtig für PyGaze!)
|
||||
gx = pd.Series(gx).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
gy = pd.Series(gy).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
|
||||
# xscaler = MinMaxScaler()
|
||||
# gxscale = xscaler.fit_transform(gx.values.reshape(-1, 1))
|
||||
|
||||
# yscaler = MinMaxScaler()
|
||||
# gyscale = yscaler.fit_transform(gx.values.reshape(-1, 1))
|
||||
|
||||
#print("xmax ymax", gxscale.max(), gyscale.max())
|
||||
|
||||
#out = np.column_stack((gxscale, gyscale))
|
||||
out = np.column_stack((gx, gy))
|
||||
|
||||
print("→ extract_gaze_signal(): Ausgabegröße:", out.shape)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def extract_pupil(df):
|
||||
"""Extrahiert Pupillengröße (beide Augen gemittelt)."""
|
||||
|
||||
pl = df["LEFT_PUPIL_DIAMETER"].replace([np.inf, -np.inf], np.nan)
|
||||
pr = df["RIGHT_PUPIL_DIAMETER"].replace([np.inf, -np.inf], np.nan)
|
||||
|
||||
vl = df.get("LEFT_PUPIL_VALIDITY")
|
||||
vr = df.get("RIGHT_PUPIL_VALIDITY")
|
||||
|
||||
if vl is None or vr is None:
|
||||
# Falls Validity-Spalten nicht vorhanden sind, versuchen wir grobe Heuristik:
|
||||
# gültig, wenn Pupillendurchmesser nicht NaN.
|
||||
validity = (~pl.isna() | ~pr.isna()).astype(int).to_numpy()
|
||||
else:
|
||||
# Falls vorhanden: 1 wenn mindestens eines der Augen gültig ist
|
||||
validity = ( (vl == 1) | (vr == 1) ).astype(int).to_numpy()
|
||||
|
||||
# Mittelwert der verfügbaren Pupillen
|
||||
p = np.mean(np.column_stack([pl, pr]), axis=1)
|
||||
|
||||
# INF/NaN reparieren
|
||||
p = pd.Series(p).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
p = p.to_numpy()
|
||||
|
||||
print("→ extract_pupil(): Pupillensignal Länge:", len(p))
|
||||
return p, validity
|
||||
|
||||
|
||||
def detect_blinks(pupil_validity, min_duration=5):
|
||||
"""Erkennt Blinks: Validity=0 → Blink."""
|
||||
blinks = []
|
||||
start = None
|
||||
|
||||
for i, v in enumerate(pupil_validity):
|
||||
if v == 0 and start is None:
|
||||
start = i
|
||||
elif v == 1 and start is not None:
|
||||
if i - start >= min_duration:
|
||||
blinks.append([start, i])
|
||||
start = None
|
||||
|
||||
return blinks
|
||||
|
||||
|
||||
def compute_IPA(pupil, fs=250):
|
||||
"""
|
||||
IPA = Index of Pupillary Activity (nach Duchowski 2018).
|
||||
Hochfrequenzanteile der Pupillenzeitreihe.
|
||||
"""
|
||||
f, Pxx = welch(pupil, fs=fs, nperseg=int(fs*2)) # 2 Sekunden Fenster
|
||||
|
||||
hf_band = (f >= 0.6) & (f <= 2.0)
|
||||
ipa = np.sum(Pxx[hf_band])
|
||||
|
||||
return ipa
|
||||
|
||||
|
||||
##############################################################################
|
||||
# 2. FEATURE-EXTRAKTION (HAUPTFUNKTION)
|
||||
##############################################################################
|
||||
|
||||
def extract_eye_features(df, window_length_sec=50, fs=250):
|
||||
"""
|
||||
df = Tobii DataFrame
|
||||
window_length_sec = Fenstergröße (z.B. W=1s)
|
||||
"""
|
||||
|
||||
print("→ extract_eye_features(): Starte Feature-Berechnung...")
|
||||
print(" Fensterlänge W =", window_length_sec, "s")
|
||||
|
||||
W = int(window_length_sec * fs) # Window größe in Samples
|
||||
|
||||
# Gaze
|
||||
gaze = extract_gaze_signal(df)
|
||||
gx, gy = gaze[:, 0], gaze[:, 1]
|
||||
print("Gültige Werte (gx):", np.sum(~np.isnan(gx)), "von", len(gx))
|
||||
print("Range:", np.nanmin(gx), np.nanmax(gx))
|
||||
print("Gültige Werte (gy):", np.sum(~np.isnan(gy)), "von", len(gy))
|
||||
print("Range:", np.nanmin(gy), np.nanmax(gy))
|
||||
|
||||
# Pupille
|
||||
pupil, pupil_validity = extract_pupil(df)
|
||||
|
||||
features = []
|
||||
|
||||
# Sliding windows
|
||||
for start in range(0, len(df), W):
|
||||
end = start + W
|
||||
if end > len(df):
|
||||
break #das letzte Fenster wird ignoriert
|
||||
|
||||
|
||||
w_gaze = gaze[start:end]
|
||||
w_pupil = pupil[start:end]
|
||||
w_valid = pupil_validity[start:end]
|
||||
|
||||
# ----------------------------
|
||||
# FIXATIONS (PyGaze)
|
||||
# ----------------------------
|
||||
time_ms = np.arange(W) * 1000.0 / fs
|
||||
|
||||
# print("gx im Fenster:", w_gaze[:,0][:20])
|
||||
# print("gy im Fenster:", w_gaze[:,1][:20])
|
||||
# print("gx diff:", np.mean(np.abs(np.diff(w_gaze[:,0]))))
|
||||
|
||||
# print("Werte X im Fenster:", w_gaze[:,0])
|
||||
# print("Werte Y im Fenster:", w_gaze[:,1])
|
||||
# print("X-Stats: min/max/diff", np.nanmin(w_gaze[:,0]), np.nanmax(w_gaze[:,0]), np.nanmean(np.abs(np.diff(w_gaze[:,0]))))
|
||||
# print("Y-Stats: min/max/diff", np.nanmin(w_gaze[:,1]), np.nanmax(w_gaze[:,1]), np.nanmean(np.abs(np.diff(w_gaze[:,1]))))
|
||||
print("time_ms:", time_ms)
|
||||
|
||||
fix, efix = fixation_detection(
|
||||
x=w_gaze[:, 0], y=w_gaze[:, 1], time=time_ms,
|
||||
missing=0.0, maxdist=0.003, mindur=10 # mindur=100ms
|
||||
)
|
||||
|
||||
#print("Raw Fixation Output:", efix[0])
|
||||
|
||||
if start == 0:
|
||||
print("DEBUG fix raw:", fix[:10])
|
||||
|
||||
# Robust fixations: PyGaze may return malformed entries
|
||||
fixation_durations = []
|
||||
for f in efix:
|
||||
print("Efix:", f[2])
|
||||
# start_t = f[1] # in ms
|
||||
# end_t = f[2] # in ms
|
||||
# duration = (end_t - start_t) / 1000.0 # in Sekunden
|
||||
|
||||
#duration = f[2] / 1000.0
|
||||
if np.isfinite(f[2]) and f[2] > 0:
|
||||
fixation_durations.append(f[2])
|
||||
|
||||
# Kategorien laut Paper
|
||||
F_short = sum(66 <= d <= 150 for d in fixation_durations)
|
||||
F_medium = sum(300 <= d <= 500 for d in fixation_durations)
|
||||
F_long = sum(d >= 1000 for d in fixation_durations)
|
||||
F_hundred = sum(d > 100 for d in fixation_durations)
|
||||
F_Cancel = sum(66 < d for d in fixation_durations)
|
||||
|
||||
# ----------------------------
|
||||
# SACCADES
|
||||
# ----------------------------
|
||||
sac, esac = saccade_detection(
|
||||
x=w_gaze[:, 0], y=w_gaze[:, 1], time=time_ms, missing=0, minlen=12, maxvel=0.2, maxacc=1
|
||||
)
|
||||
|
||||
sac_durations = [s[2] for s in esac]
|
||||
sac_amplitudes = [((s[5]-s[3])**2 + (s[6]-s[4])**2)**0.5 for s in esac]
|
||||
|
||||
# ----------------------------
|
||||
# BLINKS
|
||||
# ----------------------------
|
||||
blinks = detect_blinks(w_valid)
|
||||
blink_durations = [(b[1] - b[0]) / fs for b in blinks]
|
||||
|
||||
# ----------------------------
|
||||
# PUPIL
|
||||
# ----------------------------
|
||||
if np.all(np.isnan(w_pupil)):
|
||||
mean_pupil = np.nan
|
||||
ipa = np.nan
|
||||
else:
|
||||
mean_pupil = np.nanmean(w_pupil)
|
||||
ipa = compute_IPA(w_pupil, fs=fs)
|
||||
|
||||
# ----------------------------
|
||||
# FEATURE-TABELLE FÜLLEN
|
||||
# ----------------------------
|
||||
features.append({
|
||||
"Fix_count_short_66_150": F_short,
|
||||
"Fix_count_medium_300_500": F_medium,
|
||||
"Fix_count_long_gt_1000": F_long,
|
||||
"Fix_count_100": F_hundred,
|
||||
"Fix_cancel": F_Cancel,
|
||||
"Fix_mean_duration": np.mean(fixation_durations) if fixation_durations else 0,
|
||||
"Fix_median_duration": np.median(fixation_durations) if fixation_durations else 0,
|
||||
|
||||
"Sac_count": len(sac),
|
||||
"Sac_mean_amp": np.mean(sac_amplitudes) if sac_amplitudes else 0,
|
||||
"Sac_mean_dur": np.mean(sac_durations) if sac_durations else 0,
|
||||
"Sac_median_dur": np.median(sac_durations) if sac_durations else 0,
|
||||
|
||||
"Blink_count": len(blinks),
|
||||
"Blink_mean_dur": np.mean(blink_durations) if blink_durations else 0,
|
||||
"Blink_median_dur": np.median(blink_durations) if blink_durations else 0,
|
||||
|
||||
"Pupil_mean": mean_pupil,
|
||||
"Pupil_IPA": ipa
|
||||
})
|
||||
|
||||
|
||||
result = pd.DataFrame(features)
|
||||
print("→ extract_eye_features(): Fertig! Ergebnisgröße:", result.shape)
|
||||
|
||||
return result
|
||||
|
||||
##############################################################################
|
||||
# 3. MAIN FUNKTION
|
||||
##############################################################################
|
||||
|
||||
def main():
|
||||
print("### STARTE FEATURE-EXTRAKTION ###")
|
||||
print("Aktueller Arbeitsordner:", os.getcwd())
|
||||
|
||||
#df = pd.read_hdf("tmp22.h5", "SIGNALS", mode="r")
|
||||
df = pd.read_parquet("cleaned_0001.parquet")
|
||||
print("DataFrame geladen:", df.shape)
|
||||
|
||||
# Nur Eye-Tracking auswählen
|
||||
#eye_cols = [c for c in df.columns if "EYE_" in c]
|
||||
#df_eye = df[eye_cols]
|
||||
|
||||
#print("Eye-Tracking-Spalten:", len(eye_cols))
|
||||
#print("→", eye_cols[:10], " ...")
|
||||
|
||||
print("Reinige Eyetracking-Daten ...")
|
||||
df_eye = clean_eye_df(df)
|
||||
|
||||
# Feature Extraction
|
||||
features = extract_eye_features(df_eye, window_length_sec=50, fs=250)
|
||||
|
||||
print("\n### FEATURE-MATRIX (HEAD) ###")
|
||||
print(features.head())
|
||||
|
||||
print("\nSpeichere Output in features.csv ...")
|
||||
features.to_csv("features4.csv", index=False)
|
||||
|
||||
print("FERTIG!")
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,441 +0,0 @@
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import h5py
|
||||
import yaml
|
||||
import os
|
||||
from pathlib import Path
|
||||
from sklearn.preprocessing import MinMaxScaler
|
||||
from scipy.signal import welch
|
||||
from pygazeanalyser.detectors import fixation_detection, saccade_detection
|
||||
|
||||
|
||||
##############################################################################
|
||||
# KONFIGURATION - HIER ANPASSEN!
|
||||
##############################################################################
|
||||
INPUT_DIR = Path(r"/home/jovyan/data-paulusjafahrsimulator-gpu/new_ET_Parquet_files/")
|
||||
OUTPUT_FILE = Path(r"/home/jovyan/data-paulusjafahrsimulator-gpu/Eye_dataset_old/new_eye_dataset.parquet")
|
||||
|
||||
WINDOW_SIZE_SAMPLES = 12500 # Anzahl Samples pro Window (z.B. 1250 = 50s bei 25Hz, oder 5s bei 250Hz)
|
||||
STEP_SIZE_SAMPLES = 1250 # Schrittweite (z.B. 125 = 5s bei 25Hz, oder 0.5s bei 250Hz)
|
||||
SAMPLING_RATE = 250 # Hz
|
||||
|
||||
|
||||
##############################################################################
|
||||
# 1. HELFERFUNKTIONEN
|
||||
##############################################################################
|
||||
def clean_eye_df(df):
|
||||
"""
|
||||
Entfernt alle Zeilen, die keine echten Eyetracking-Daten enthalten.
|
||||
Löst das Problem, dass das Haupt-DataFrame NaN-Zeilen für andere Sensoren enthält.
|
||||
"""
|
||||
eye_cols = [c for c in df.columns if c.startswith("EYE_")]
|
||||
df_eye = df[eye_cols]
|
||||
|
||||
# INF → NaN
|
||||
df_eye = df_eye.replace([np.inf, -np.inf], np.nan)
|
||||
|
||||
# Nur Zeilen behalten, wo es echte Eyetracking-Daten gibt
|
||||
df_eye = df_eye.dropna(subset=eye_cols, how="all")
|
||||
|
||||
print(f" Eyetracking-Zeilen: {len(df)} → {len(df_eye)}")
|
||||
|
||||
return df_eye.reset_index(drop=True)
|
||||
|
||||
|
||||
def extract_gaze_signal(df):
|
||||
"""
|
||||
Extrahiert 2D-Gaze-Positionen auf dem Display,
|
||||
maskiert ungültige Samples und interpoliert Lücken.
|
||||
"""
|
||||
# Gaze-Spalten
|
||||
gx_L = df["EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_X"].astype(float).copy()
|
||||
gy_L = df["EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_Y"].astype(float).copy()
|
||||
gx_R = df["EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_X"].astype(float).copy()
|
||||
gy_R = df["EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_Y"].astype(float).copy()
|
||||
|
||||
# Validity-Spalten (1 = gültig)
|
||||
val_L = (df["EYE_LEFT_GAZE_POINT_VALIDITY"] == 1)
|
||||
val_R = (df["EYE_RIGHT_GAZE_POINT_VALIDITY"] == 1)
|
||||
|
||||
# Inf ersetzen mit NaN (kommt bei Tobii bei Blinks vor)
|
||||
gx_L.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
gy_L.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
gx_R.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
gy_R.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
|
||||
# Ungültige Werte maskieren
|
||||
gx_L[~val_L] = np.nan
|
||||
gy_L[~val_L] = np.nan
|
||||
gx_R[~val_R] = np.nan
|
||||
gy_R[~val_R] = np.nan
|
||||
|
||||
# Mittelwert der beiden Augen pro Sample (nanmean ist robust)
|
||||
gx = np.mean(np.column_stack([gx_L, gx_R]), axis=1)
|
||||
gy = np.mean(np.column_stack([gy_L, gy_R]), axis=1)
|
||||
|
||||
# Interpolation (wichtig für PyGaze!)
|
||||
gx = pd.Series(gx).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
gy = pd.Series(gy).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
|
||||
xscaler = MinMaxScaler()
|
||||
gxscale = xscaler.fit_transform(gx.values.reshape(-1, 1))
|
||||
|
||||
yscaler = MinMaxScaler()
|
||||
gyscale = yscaler.fit_transform(gy.values.reshape(-1, 1))
|
||||
|
||||
out = np.column_stack((gxscale, gyscale))
|
||||
return out
|
||||
|
||||
|
||||
def extract_pupil(df):
|
||||
"""Extrahiert Pupillengröße (beide Augen gemittelt)."""
|
||||
pl = df["EYE_LEFT_PUPIL_DIAMETER"].replace([np.inf, -np.inf], np.nan)
|
||||
pr = df["EYE_RIGHT_PUPIL_DIAMETER"].replace([np.inf, -np.inf], np.nan)
|
||||
|
||||
vl = df.get("EYE_LEFT_PUPIL_VALIDITY")
|
||||
vr = df.get("EYE_RIGHT_PUPIL_VALIDITY")
|
||||
|
||||
if vl is None or vr is None:
|
||||
validity = (~pl.isna() | ~pr.isna()).astype(int).to_numpy()
|
||||
else:
|
||||
validity = ((vl == 1) | (vr == 1)).astype(int).to_numpy()
|
||||
|
||||
# Mittelwert der verfügbaren Pupillen
|
||||
p = np.mean(np.column_stack([pl, pr]), axis=1)
|
||||
|
||||
# INF/NaN reparieren
|
||||
p = pd.Series(p).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
p = p.to_numpy()
|
||||
|
||||
return p, validity
|
||||
|
||||
|
||||
def detect_blinks(pupil_validity, min_duration=5):
|
||||
"""Erkennt Blinks: Validity=0 → Blink."""
|
||||
blinks = []
|
||||
start = None
|
||||
|
||||
for i, v in enumerate(pupil_validity):
|
||||
if v == 0 and start is None:
|
||||
start = i
|
||||
elif v == 1 and start is not None:
|
||||
if i - start >= min_duration:
|
||||
blinks.append([start, i])
|
||||
start = None
|
||||
|
||||
return blinks
|
||||
|
||||
|
||||
def compute_IPA(pupil, fs=250):
|
||||
"""
|
||||
IPA = Index of Pupillary Activity (nach Duchowski 2018).
|
||||
Hochfrequenzanteile der Pupillenzeitreihe.
|
||||
"""
|
||||
f, Pxx = welch(pupil, fs=fs, nperseg=int(fs*2)) # 2 Sekunden Fenster
|
||||
|
||||
hf_band = (f >= 0.6) & (f <= 2.0)
|
||||
ipa = np.sum(Pxx[hf_band])
|
||||
|
||||
return ipa
|
||||
|
||||
|
||||
##############################################################################
|
||||
# 2. FEATURE-EXTRAKTION MIT SLIDING WINDOW
|
||||
##############################################################################
|
||||
|
||||
def extract_eye_features_sliding(df_eye, df_meta, window_size, step_size, fs=250):
|
||||
"""
|
||||
Extrahiert Features mit Sliding Window aus einem einzelnen Level/Phase.
|
||||
|
||||
Parameters:
|
||||
-----------
|
||||
df_eye : DataFrame
|
||||
Eye-Tracking Daten (bereits gereinigt)
|
||||
df_meta : DataFrame
|
||||
Metadaten (subjectID, rowID, STUDY, LEVEL, PHASE)
|
||||
window_size : int
|
||||
Anzahl Samples pro Window
|
||||
step_size : int
|
||||
Schrittweite in Samples
|
||||
fs : int
|
||||
Sampling Rate in Hz
|
||||
"""
|
||||
# Gaze
|
||||
gaze = extract_gaze_signal(df_eye)
|
||||
|
||||
# Pupille
|
||||
pupil, pupil_validity = extract_pupil(df_eye)
|
||||
|
||||
features = []
|
||||
num_windows = (len(df_eye) - window_size) // step_size + 1
|
||||
|
||||
if num_windows <= 0:
|
||||
return pd.DataFrame()
|
||||
|
||||
for i in range(num_windows):
|
||||
start_idx = i * step_size
|
||||
end_idx = start_idx + window_size
|
||||
|
||||
w_gaze = gaze[start_idx:end_idx]
|
||||
w_pupil = pupil[start_idx:end_idx]
|
||||
w_valid = pupil_validity[start_idx:end_idx]
|
||||
|
||||
# Metadaten für dieses Window
|
||||
meta_row = df_meta.iloc[start_idx]
|
||||
|
||||
# ----------------------------
|
||||
# FIXATIONS (PyGaze)
|
||||
# ----------------------------
|
||||
time_ms = np.arange(window_size) * 1000.0 / fs
|
||||
|
||||
fix, efix = fixation_detection(
|
||||
x=w_gaze[:, 0], y=w_gaze[:, 1], time=time_ms,
|
||||
missing=0.0, maxdist=0.003, mindur=10
|
||||
)
|
||||
|
||||
fixation_durations = []
|
||||
for f in efix:
|
||||
if np.isfinite(f[2]) and f[2] > 0:
|
||||
fixation_durations.append(f[2])
|
||||
|
||||
# Kategorien laut Paper
|
||||
F_short = sum(66 <= d <= 150 for d in fixation_durations)
|
||||
F_medium = sum(300 <= d <= 500 for d in fixation_durations)
|
||||
F_long = sum(d >= 1000 for d in fixation_durations)
|
||||
F_hundred = sum(d > 100 for d in fixation_durations)
|
||||
# F_Cancel = sum(66 < d for d in fixation_durations)
|
||||
|
||||
# ----------------------------
|
||||
# SACCADES
|
||||
# ----------------------------
|
||||
sac, esac = saccade_detection(
|
||||
x=w_gaze[:, 0], y=w_gaze[:, 1], time=time_ms,
|
||||
missing=0, minlen=12, maxvel=0.2, maxacc=1
|
||||
)
|
||||
|
||||
sac_durations = [s[2] for s in esac]
|
||||
sac_amplitudes = [((s[5]-s[3])**2 + (s[6]-s[4])**2)**0.5 for s in esac]
|
||||
|
||||
# ----------------------------
|
||||
# BLINKS
|
||||
# ----------------------------
|
||||
blinks = detect_blinks(w_valid)
|
||||
blink_durations = [(b[1] - b[0]) / fs for b in blinks]
|
||||
|
||||
# ----------------------------
|
||||
# PUPIL
|
||||
# ----------------------------
|
||||
if np.all(np.isnan(w_pupil)):
|
||||
mean_pupil = np.nan
|
||||
ipa = np.nan
|
||||
else:
|
||||
mean_pupil = np.nanmean(w_pupil)
|
||||
ipa = compute_IPA(w_pupil, fs=fs)
|
||||
|
||||
# ----------------------------
|
||||
# FEATURE-DICTIONARY
|
||||
# ----------------------------
|
||||
features.append({
|
||||
# Metadaten
|
||||
'subjectID': meta_row['subjectID'],
|
||||
'start_time': meta_row['rowID'],
|
||||
'STUDY': meta_row.get('STUDY', np.nan),
|
||||
'LEVEL': meta_row.get('LEVEL', np.nan),
|
||||
'PHASE': meta_row.get('PHASE', np.nan),
|
||||
|
||||
# Fixation Features
|
||||
"Fix_count_short_66_150": F_short,
|
||||
"Fix_count_medium_300_500": F_medium,
|
||||
"Fix_count_long_gt_1000": F_long,
|
||||
"Fix_count_100": F_hundred,
|
||||
# "Fix_cancel": F_Cancel,
|
||||
"Fix_mean_duration": np.mean(fixation_durations) if fixation_durations else 0,
|
||||
"Fix_median_duration": np.median(fixation_durations) if fixation_durations else 0,
|
||||
|
||||
# Saccade Features
|
||||
"Sac_count": len(sac),
|
||||
"Sac_mean_amp": np.mean(sac_amplitudes) if sac_amplitudes else 0,
|
||||
"Sac_mean_dur": np.mean(sac_durations) if sac_durations else 0,
|
||||
"Sac_median_dur": np.median(sac_durations) if sac_durations else 0,
|
||||
|
||||
# Blink Features
|
||||
"Blink_count": len(blinks),
|
||||
"Blink_mean_dur": np.mean(blink_durations) if blink_durations else 0,
|
||||
"Blink_median_dur": np.median(blink_durations) if blink_durations else 0,
|
||||
|
||||
# Pupil Features
|
||||
"Pupil_mean": mean_pupil,
|
||||
"Pupil_IPA": ipa
|
||||
})
|
||||
|
||||
return pd.DataFrame(features)
|
||||
|
||||
|
||||
##############################################################################
|
||||
# 3. BATCH-VERARBEITUNG
|
||||
##############################################################################
|
||||
|
||||
def process_parquet_directory(input_dir, output_file, window_size, step_size, fs=250):
|
||||
"""
|
||||
Verarbeitet alle Parquet-Dateien in einem Verzeichnis.
|
||||
|
||||
Parameters:
|
||||
-----------
|
||||
input_dir : str
|
||||
Pfad zum Verzeichnis mit Parquet-Dateien
|
||||
output_file : str
|
||||
Pfad für die Ausgabe-Parquet-Datei
|
||||
window_size : int
|
||||
Window-Größe in Samples
|
||||
step_size : int
|
||||
Schrittweite in Samples
|
||||
fs : int
|
||||
Sampling Rate in Hz
|
||||
"""
|
||||
input_path = Path(input_dir)
|
||||
parquet_files = sorted(input_path.glob("*.parquet"))
|
||||
|
||||
if not parquet_files:
|
||||
print(f"FEHLER: Keine Parquet-Dateien in {input_dir} gefunden!")
|
||||
return
|
||||
|
||||
print(f"\n{'='*70}")
|
||||
print(f"STARTE BATCH-VERARBEITUNG")
|
||||
print(f"{'='*70}")
|
||||
print(f"Gefundene Dateien: {len(parquet_files)}")
|
||||
print(f"Window Size: {window_size} Samples ({window_size/fs:.1f}s bei {fs}Hz)")
|
||||
print(f"Step Size: {step_size} Samples ({step_size/fs:.1f}s bei {fs}Hz)")
|
||||
print(f"{'='*70}\n")
|
||||
|
||||
all_features = []
|
||||
|
||||
for file_idx, parquet_file in enumerate(parquet_files, 1):
|
||||
print(f"\n[{file_idx}/{len(parquet_files)}] Verarbeite: {parquet_file.name}")
|
||||
|
||||
try:
|
||||
# Lade Parquet-Datei
|
||||
df = pd.read_parquet(parquet_file)
|
||||
print(f" Einträge geladen: {len(df)}")
|
||||
|
||||
# Prüfe ob benötigte Spalten vorhanden sind
|
||||
required_cols = ['subjectID', 'rowID']
|
||||
missing_cols = [col for col in required_cols if col not in df.columns]
|
||||
if missing_cols:
|
||||
print(f" WARNUNG: Fehlende Spalten: {missing_cols} - Überspringe Datei")
|
||||
continue
|
||||
|
||||
# Reinige Eye-Tracking-Daten
|
||||
df_eye = clean_eye_df(df)
|
||||
|
||||
if len(df_eye) == 0:
|
||||
print(f" WARNUNG: Keine gültigen Eye-Tracking-Daten - Überspringe Datei")
|
||||
continue
|
||||
|
||||
# Metadaten extrahieren (aligned mit df_eye)
|
||||
meta_cols = ['subjectID', 'rowID']
|
||||
if 'STUDY' in df.columns:
|
||||
meta_cols.append('STUDY')
|
||||
if 'LEVEL' in df.columns:
|
||||
meta_cols.append('LEVEL')
|
||||
if 'PHASE' in df.columns:
|
||||
meta_cols.append('PHASE')
|
||||
|
||||
df_meta = df[meta_cols].iloc[df_eye.index].reset_index(drop=True)
|
||||
|
||||
# Gruppiere nach STUDY, LEVEL, PHASE (falls vorhanden)
|
||||
group_cols = [col for col in ['STUDY', 'LEVEL', 'PHASE'] if col in df_meta.columns]
|
||||
|
||||
if group_cols:
|
||||
print(f" Gruppiere nach: {', '.join(group_cols)}")
|
||||
for group_vals, group_df in df_meta.groupby(group_cols, sort=False):
|
||||
group_eye = df_eye.iloc[group_df.index].reset_index(drop=True)
|
||||
group_meta = group_df.reset_index(drop=True)
|
||||
|
||||
print(f" Gruppe {group_vals}: {len(group_eye)} Samples", end=" → ")
|
||||
|
||||
features_df = extract_eye_features_sliding(
|
||||
group_eye, group_meta, window_size, step_size, fs
|
||||
)
|
||||
|
||||
if not features_df.empty:
|
||||
all_features.append(features_df)
|
||||
print(f"{len(features_df)} Windows")
|
||||
else:
|
||||
print("Zu wenige Daten")
|
||||
else:
|
||||
# Keine Gruppierung
|
||||
print(f" Keine Gruppierungsspalten gefunden")
|
||||
features_df = extract_eye_features_sliding(
|
||||
df_eye, df_meta, window_size, step_size, fs
|
||||
)
|
||||
|
||||
if not features_df.empty:
|
||||
all_features.append(features_df)
|
||||
print(f" → {len(features_df)} Windows erstellt")
|
||||
else:
|
||||
print(f" → Zu wenige Daten")
|
||||
|
||||
except Exception as e:
|
||||
print(f" FEHLER bei Verarbeitung: {str(e)}")
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
continue
|
||||
|
||||
# Kombiniere alle Features
|
||||
if not all_features:
|
||||
print("\nKEINE FEATURES EXTRAHIERT!")
|
||||
return None
|
||||
|
||||
print(f"\n{'='*70}")
|
||||
print(f"ZUSAMMENFASSUNG")
|
||||
print(f"{'='*70}")
|
||||
|
||||
final_df = pd.concat(all_features, ignore_index=True)
|
||||
|
||||
print(f"Gesamt Windows: {len(final_df)}")
|
||||
print(f"Spalten: {len(final_df.columns)}")
|
||||
print(f"Subjects: {final_df['subjectID'].nunique()}")
|
||||
|
||||
# Speichere Ergebnis
|
||||
output_path = Path(output_file)
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
final_df.to_parquet(output_file, index=False)
|
||||
|
||||
print(f"\n✓ Ergebnis gespeichert: {output_file}")
|
||||
print(f"{'='*70}\n")
|
||||
|
||||
return final_df
|
||||
|
||||
|
||||
##############################################################################
|
||||
# 4. MAIN
|
||||
##############################################################################
|
||||
|
||||
def main():
|
||||
print("\n" + "="*70)
|
||||
print("EYE-TRACKING FEATURE EXTRAKTION - BATCH MODE")
|
||||
print("="*70)
|
||||
|
||||
result = process_parquet_directory(
|
||||
input_dir=INPUT_DIR,
|
||||
output_file=OUTPUT_FILE,
|
||||
window_size=WINDOW_SIZE_SAMPLES,
|
||||
step_size=STEP_SIZE_SAMPLES,
|
||||
fs=SAMPLING_RATE
|
||||
)
|
||||
|
||||
if result is not None:
|
||||
print("\nErste 5 Zeilen des Ergebnisses:")
|
||||
print(result.head())
|
||||
|
||||
print("\nSpalten-Übersicht:")
|
||||
print(result.columns.tolist())
|
||||
|
||||
print("\nDatentypen:")
|
||||
print(result.dtypes)
|
||||
|
||||
print("\n✓ FERTIG!\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,323 +0,0 @@
|
||||
import numpy as np
|
||||
import pandas as pd
|
||||
import h5py
|
||||
import yaml
|
||||
import owncloud
|
||||
import os
|
||||
from sklearn.preprocessing import MinMaxScaler
|
||||
from scipy.signal import welch
|
||||
from pygazeanalyser.detectors import fixation_detection, saccade_detection
|
||||
|
||||
|
||||
##############################################################################
|
||||
# 1. HELFERFUNKTIONEN
|
||||
##############################################################################
|
||||
def clean_eye_df(df):
|
||||
"""
|
||||
Entfernt alle Zeilen, die keine echten Eyetracking-Daten enthalten.
|
||||
Löst das Problem, dass das Haupt-DataFrame NaN-Zeilen für andere Sensoren enthält.
|
||||
"""
|
||||
eye_cols = [c for c in df.columns if "EYE_" in c]
|
||||
df_eye = df[eye_cols]
|
||||
|
||||
# INF → NaN
|
||||
df_eye = df_eye.replace([np.inf, -np.inf], np.nan)
|
||||
|
||||
# Nur Zeilen behalten, wo es echte Eyetracking-Daten gibt
|
||||
df_eye = df_eye.dropna(subset=eye_cols, how="all")
|
||||
|
||||
print("Eyetracking-Zeilen vorher:", len(df))
|
||||
print("Eyetracking-Zeilen nachher:", len(df_eye))
|
||||
|
||||
#Index zurücksetzen
|
||||
return df_eye.reset_index(drop=True)
|
||||
|
||||
|
||||
def extract_gaze_signal(df):
|
||||
"""
|
||||
Extrahiert 2D-Gaze-Positionen auf dem Display,
|
||||
maskiert ungültige Samples und interpoliert Lücken.
|
||||
"""
|
||||
|
||||
print("→ extract_gaze_signal(): Eingabegröße:", df.shape)
|
||||
|
||||
# Gaze-Spalten
|
||||
gx_L = df["EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_X"].astype(float).copy()
|
||||
gy_L = df["EYE_LEFT_GAZE_POINT_ON_DISPLAY_AREA_Y"].astype(float).copy()
|
||||
gx_R = df["EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_X"].astype(float).copy()
|
||||
gy_R = df["EYE_RIGHT_GAZE_POINT_ON_DISPLAY_AREA_Y"].astype(float).copy()
|
||||
|
||||
|
||||
# Validity-Spalten (1 = gültig)
|
||||
val_L = (df["EYE_LEFT_GAZE_POINT_VALIDITY"] == 1)
|
||||
val_R = (df["EYE_RIGHT_GAZE_POINT_VALIDITY"] == 1)
|
||||
|
||||
# Inf ersetzen mit NaN (kommt bei Tobii bei Blinks vor)
|
||||
gx_L.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
gy_L.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
gx_R.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
gy_R.replace([np.inf, -np.inf], np.nan, inplace=True)
|
||||
|
||||
# Ungültige Werte maskieren
|
||||
gx_L[~val_L] = np.nan
|
||||
gy_L[~val_L] = np.nan
|
||||
gx_R[~val_R] = np.nan
|
||||
gy_R[~val_R] = np.nan
|
||||
|
||||
# Mittelwert der beiden Augen pro Sample (nanmean ist robust)
|
||||
gx = np.mean(np.column_stack([gx_L, gx_R]), axis=1)
|
||||
gy = np.mean(np.column_stack([gy_L, gy_R]), axis=1)
|
||||
|
||||
# Interpolation (wichtig für PyGaze!)
|
||||
gx = pd.Series(gx).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
gy = pd.Series(gy).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
|
||||
xscaler = MinMaxScaler()
|
||||
gxscale = xscaler.fit_transform(gx.values.reshape(-1, 1))
|
||||
|
||||
yscaler = MinMaxScaler()
|
||||
gyscale = yscaler.fit_transform(gx.values.reshape(-1, 1))
|
||||
|
||||
print("xmax ymax", gxscale.max(), gyscale.max())
|
||||
|
||||
out = np.column_stack((gxscale, gyscale))
|
||||
|
||||
print("→ extract_gaze_signal(): Ausgabegröße:", out.shape)
|
||||
|
||||
return out
|
||||
|
||||
|
||||
def extract_pupil(df):
|
||||
"""Extrahiert Pupillengröße (beide Augen gemittelt)."""
|
||||
|
||||
pl = df["EYE_LEFT_PUPIL_DIAMETER"].replace([np.inf, -np.inf], np.nan)
|
||||
pr = df["EYE_RIGHT_PUPIL_DIAMETER"].replace([np.inf, -np.inf], np.nan)
|
||||
|
||||
vl = df.get("EYE_LEFT_PUPIL_VALIDITY")
|
||||
vr = df.get("EYE_RIGHT_PUPIL_VALIDITY")
|
||||
|
||||
if vl is None or vr is None:
|
||||
# Falls Validity-Spalten nicht vorhanden sind, versuchen wir grobe Heuristik:
|
||||
# gültig, wenn Pupillendurchmesser nicht NaN.
|
||||
validity = (~pl.isna() | ~pr.isna()).astype(int).to_numpy()
|
||||
else:
|
||||
# Falls vorhanden: 1 wenn mindestens eines der Augen gültig ist
|
||||
validity = ( (vl == 1) | (vr == 1) ).astype(int).to_numpy()
|
||||
|
||||
# Mittelwert der verfügbaren Pupillen
|
||||
p = np.mean(np.column_stack([pl, pr]), axis=1)
|
||||
|
||||
# INF/NaN reparieren
|
||||
p = pd.Series(p).interpolate(limit=50, limit_direction="both").bfill().ffill()
|
||||
p = p.to_numpy()
|
||||
|
||||
print("→ extract_pupil(): Pupillensignal Länge:", len(p))
|
||||
return p, validity
|
||||
|
||||
|
||||
def detect_blinks(pupil_validity, min_duration=5):
|
||||
"""Erkennt Blinks: Validity=0 → Blink."""
|
||||
blinks = []
|
||||
start = None
|
||||
|
||||
for i, v in enumerate(pupil_validity):
|
||||
if v == 0 and start is None:
|
||||
start = i
|
||||
elif v == 1 and start is not None:
|
||||
if i - start >= min_duration:
|
||||
blinks.append([start, i])
|
||||
start = None
|
||||
|
||||
return blinks
|
||||
|
||||
|
||||
def compute_IPA(pupil, fs=250):
|
||||
"""
|
||||
IPA = Index of Pupillary Activity (nach Duchowski 2018).
|
||||
Hochfrequenzanteile der Pupillenzeitreihe.
|
||||
"""
|
||||
f, Pxx = welch(pupil, fs=fs, nperseg=int(fs*2)) # 2 Sekunden Fenster
|
||||
|
||||
hf_band = (f >= 0.6) & (f <= 2.0)
|
||||
ipa = np.sum(Pxx[hf_band])
|
||||
|
||||
return ipa
|
||||
|
||||
|
||||
##############################################################################
|
||||
# 2. FEATURE-EXTRAKTION (HAUPTFUNKTION)
|
||||
##############################################################################
|
||||
|
||||
def extract_eye_features(df, window_length_sec=50, fs=250):
|
||||
"""
|
||||
df = Tobii DataFrame
|
||||
window_length_sec = Fenstergröße (z.B. W=1s)
|
||||
"""
|
||||
|
||||
print("→ extract_eye_features(): Starte Feature-Berechnung...")
|
||||
print(" Fensterlänge W =", window_length_sec, "s")
|
||||
|
||||
W = int(window_length_sec * fs) # Window größe in Samples
|
||||
|
||||
# Gaze
|
||||
gaze = extract_gaze_signal(df)
|
||||
gx, gy = gaze[:, 0], gaze[:, 1]
|
||||
print("Gültige Werte (gx):", np.sum(~np.isnan(gx)), "von", len(gx))
|
||||
print("Range:", np.nanmin(gx), np.nanmax(gx))
|
||||
print("Gültige Werte (gy):", np.sum(~np.isnan(gy)), "von", len(gy))
|
||||
print("Range:", np.nanmin(gy), np.nanmax(gy))
|
||||
|
||||
# Pupille
|
||||
pupil, pupil_validity = extract_pupil(df)
|
||||
|
||||
features = []
|
||||
|
||||
# Sliding windows
|
||||
for start in range(0, len(df), W):
|
||||
end = start + W
|
||||
if end > len(df):
|
||||
break #das letzte Fenster wird ignoriert
|
||||
|
||||
|
||||
w_gaze = gaze[start:end]
|
||||
w_pupil = pupil[start:end]
|
||||
w_valid = pupil_validity[start:end]
|
||||
|
||||
# ----------------------------
|
||||
# FIXATIONS (PyGaze)
|
||||
# ----------------------------
|
||||
time_ms = np.arange(W) * 1000.0 / fs
|
||||
|
||||
# print("gx im Fenster:", w_gaze[:,0][:20])
|
||||
# print("gy im Fenster:", w_gaze[:,1][:20])
|
||||
# print("gx diff:", np.mean(np.abs(np.diff(w_gaze[:,0]))))
|
||||
|
||||
# print("Werte X im Fenster:", w_gaze[:,0])
|
||||
# print("Werte Y im Fenster:", w_gaze[:,1])
|
||||
# print("X-Stats: min/max/diff", np.nanmin(w_gaze[:,0]), np.nanmax(w_gaze[:,0]), np.nanmean(np.abs(np.diff(w_gaze[:,0]))))
|
||||
# print("Y-Stats: min/max/diff", np.nanmin(w_gaze[:,1]), np.nanmax(w_gaze[:,1]), np.nanmean(np.abs(np.diff(w_gaze[:,1]))))
|
||||
print("time_ms:", time_ms)
|
||||
|
||||
fix, efix = fixation_detection(
|
||||
x=w_gaze[:, 0], y=w_gaze[:, 1], time=time_ms,
|
||||
missing=0.0, maxdist=0.001, mindur=65 # mindur=100ms
|
||||
)
|
||||
|
||||
#print("Raw Fixation Output:", efix[0])
|
||||
|
||||
if start == 0:
|
||||
print("DEBUG fix raw:", fix[:10])
|
||||
|
||||
# Robust fixations: PyGaze may return malformed entries
|
||||
fixation_durations = []
|
||||
for f in efix:
|
||||
print("Efix:", f[2])
|
||||
# start_t = f[1] # in ms
|
||||
# end_t = f[2] # in ms
|
||||
# duration = (end_t - start_t) / 1000.0 # in Sekunden
|
||||
|
||||
#duration = f[2] / 1000.0
|
||||
if np.isfinite(f[2]) and f[2] > 0:
|
||||
fixation_durations.append(f[2])
|
||||
|
||||
# Kategorien laut Paper
|
||||
F_short = sum(66 <= d <= 150 for d in fixation_durations)
|
||||
F_medium = sum(300 <= d <= 500 for d in fixation_durations)
|
||||
F_long = sum(d >= 1000 for d in fixation_durations)
|
||||
F_hundred = sum(d > 100 for d in fixation_durations)
|
||||
F_Cancel = sum(66 < d for d in fixation_durations)
|
||||
|
||||
# ----------------------------
|
||||
# SACCADES
|
||||
# ----------------------------
|
||||
sac, esac = saccade_detection(
|
||||
x=w_gaze[:, 0], y=w_gaze[:, 1], time=time_ms, missing=0, minlen=12, maxvel=0.2, maxacc=1
|
||||
)
|
||||
|
||||
sac_durations = [s[2] for s in esac]
|
||||
sac_amplitudes = [((s[5]-s[3])**2 + (s[6]-s[4])**2)**0.5 for s in esac]
|
||||
|
||||
# ----------------------------
|
||||
# BLINKS
|
||||
# ----------------------------
|
||||
blinks = detect_blinks(w_valid)
|
||||
blink_durations = [(b[1] - b[0]) / fs for b in blinks]
|
||||
|
||||
# ----------------------------
|
||||
# PUPIL
|
||||
# ----------------------------
|
||||
if np.all(np.isnan(w_pupil)):
|
||||
mean_pupil = np.nan
|
||||
ipa = np.nan
|
||||
else:
|
||||
mean_pupil = np.nanmean(w_pupil)
|
||||
ipa = compute_IPA(w_pupil, fs=fs)
|
||||
|
||||
# ----------------------------
|
||||
# FEATURE-TABELLE FÜLLEN
|
||||
# ----------------------------
|
||||
features.append({
|
||||
"Fix_count_short_66_150": F_short,
|
||||
"Fix_count_medium_300_500": F_medium,
|
||||
"Fix_count_long_gt_1000": F_long,
|
||||
"Fix_count_100": F_hundred,
|
||||
"Fix_cancel": F_Cancel,
|
||||
"Fix_mean_duration": np.mean(fixation_durations) if fixation_durations else 0,
|
||||
"Fix_median_duration": np.median(fixation_durations) if fixation_durations else 0,
|
||||
|
||||
"Sac_count": len(sac),
|
||||
"Sac_mean_amp": np.mean(sac_amplitudes) if sac_amplitudes else 0,
|
||||
"Sac_mean_dur": np.mean(sac_durations) if sac_durations else 0,
|
||||
"Sac_median_dur": np.median(sac_durations) if sac_durations else 0,
|
||||
|
||||
"Blink_count": len(blinks),
|
||||
"Blink_mean_dur": np.mean(blink_durations) if blink_durations else 0,
|
||||
"Blink_median_dur": np.median(blink_durations) if blink_durations else 0,
|
||||
|
||||
"Pupil_mean": mean_pupil,
|
||||
"Pupil_IPA": ipa
|
||||
})
|
||||
|
||||
|
||||
result = pd.DataFrame(features)
|
||||
print("→ extract_eye_features(): Fertig! Ergebnisgröße:", result.shape)
|
||||
|
||||
return result
|
||||
|
||||
##############################################################################
|
||||
# 3. MAIN FUNKTION
|
||||
##############################################################################
|
||||
|
||||
def main():
|
||||
print("### STARTE FEATURE-EXTRAKTION ###")
|
||||
print("Aktueller Arbeitsordner:", os.getcwd())
|
||||
|
||||
df = pd.read_hdf("tmp22.h5", "SIGNALS", mode="r")
|
||||
#df = pd.read_parquet("cleaned_0001.parquet")
|
||||
print("DataFrame geladen:", df.shape)
|
||||
|
||||
# Nur Eye-Tracking auswählen
|
||||
#eye_cols = [c for c in df.columns if "EYE_" in c]
|
||||
#df_eye = df[eye_cols]
|
||||
|
||||
#print("Eye-Tracking-Spalten:", len(eye_cols))
|
||||
#print("→", eye_cols[:10], " ...")
|
||||
|
||||
print("Reinige Eyetracking-Daten ...")
|
||||
df_eye = clean_eye_df(df)
|
||||
|
||||
# Feature Extraction
|
||||
features = extract_eye_features(df_eye, window_length_sec=50, fs=250)
|
||||
|
||||
print("\n### FEATURE-MATRIX (HEAD) ###")
|
||||
print(features.head())
|
||||
|
||||
print("\nSpeichere Output in features.csv ...")
|
||||
features.to_csv("features2.csv", index=False)
|
||||
|
||||
print("FERTIG!")
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
+36
-29
@@ -1,72 +1,79 @@
|
||||
import math
|
||||
|
||||
def fixation_radius_normalized(theta_deg: float,
|
||||
|
||||
def fixation_radius_normalized(
|
||||
theta_deg: float,
|
||||
distance_cm: float,
|
||||
screen_width_cm: float,
|
||||
screen_height_cm: float,
|
||||
resolution_x: int,
|
||||
resolution_y: int,
|
||||
method: str = "max"):
|
||||
method: str = "max",
|
||||
):
|
||||
"""
|
||||
Berechnet den PyGaze-Fixationsradius für normierte Gaze-Daten in [0,1].
|
||||
Compute the PyGaze fixation radius for normalized gaze data in [0, 1].
|
||||
"""
|
||||
# Schritt 1: visueller Winkel → physische Distanz (cm)
|
||||
# Visual angle to physical distance (cm)
|
||||
delta_cm = 2 * distance_cm * math.tan(math.radians(theta_deg) / 2)
|
||||
|
||||
# Schritt 2: physische Distanz → Pixel
|
||||
# Physical distance to pixels
|
||||
delta_px_x = delta_cm * (resolution_x / screen_width_cm)
|
||||
delta_px_y = delta_cm * (resolution_y / screen_height_cm)
|
||||
|
||||
# Pixelradius
|
||||
# Pixel radius
|
||||
if method == "max":
|
||||
r_px = max(delta_px_x, delta_px_y)
|
||||
else:
|
||||
r_px = math.sqrt(delta_px_x**2 + delta_px_y**2)
|
||||
|
||||
# Schritt 3: Pixelradius → normierter Radius
|
||||
# Pixel radius to normalized radius
|
||||
r_norm_x = r_px / resolution_x
|
||||
r_norm_y = r_px / resolution_y
|
||||
|
||||
if method == "max":
|
||||
return max(r_norm_x, r_norm_y)
|
||||
else:
|
||||
return math.sqrt(r_norm_x**2 + r_norm_y**2)
|
||||
|
||||
|
||||
def run_example():
|
||||
# Example: 55" 4k monitor
|
||||
screen_width_cm = 3 * 121.8
|
||||
screen_height_cm = 68.5
|
||||
resolution_x = 3 * 3840
|
||||
resolution_y = 2160
|
||||
distance_to_screen_cm = 120
|
||||
max_angle = 1.0
|
||||
|
||||
|
||||
|
||||
|
||||
# Beispiel: 55" 4k Monitor
|
||||
screen_width_cm = 3*121.8
|
||||
screen_height_cm = 68.5
|
||||
resolution_x = 3*3840
|
||||
resolution_y = 2160
|
||||
distance_to_screen_cm = 120
|
||||
method = 'max'
|
||||
max_angle= 1.0
|
||||
|
||||
maxdist_px = fixation_radius_normalized(theta_deg=max_angle,
|
||||
maxdist_px = fixation_radius_normalized(
|
||||
theta_deg=max_angle,
|
||||
distance_cm=distance_to_screen_cm,
|
||||
screen_width_cm=screen_width_cm,
|
||||
screen_height_cm=screen_height_cm,
|
||||
resolution_x=resolution_x,
|
||||
resolution_y=resolution_y,
|
||||
method=method)
|
||||
method="max",
|
||||
)
|
||||
print("PyGaze max_dist (max):", maxdist_px)
|
||||
|
||||
print("PyGaze max_dist (max):", maxdist_px)
|
||||
|
||||
method = 'euclid'
|
||||
maxdist_px = fixation_radius_normalized(theta_deg=max_angle,
|
||||
maxdist_px = fixation_radius_normalized(
|
||||
theta_deg=max_angle,
|
||||
distance_cm=distance_to_screen_cm,
|
||||
screen_width_cm=screen_width_cm,
|
||||
screen_height_cm=screen_height_cm,
|
||||
resolution_x=resolution_x,
|
||||
resolution_y=resolution_y,
|
||||
method=method)
|
||||
method="euclid",
|
||||
)
|
||||
print("PyGaze max_dist (euclid):", maxdist_px)
|
||||
|
||||
print("PyGaze max_dist (euclid):", maxdist_px)
|
||||
|
||||
# Passt noch nicht zu der Breite
|
||||
def main():
|
||||
run_example()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
# Reference
|
||||
# https://osdoc.cogsci.nl/4.0/de/visualangle/
|
||||
# https://reference.org/facts/Visual_angle/LUw29zy7
|
||||
@@ -1,156 +0,0 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "2b3fface",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "74f1f5ec",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df= pd.read_parquet(r\"C:\\Users\\micha\\FAUbox\\WS2526_Fahrsimulator_MSY (Celina Korzer)\\AU_dataset\\output_windowed.parquet\")\n",
|
||||
"print(df.shape)\n",
|
||||
"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "05775454",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df.head()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "99e17328",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df.tail()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "69e53731",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df.info()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "3754c664",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Zeigt alle Kombinationen mit Häufigkeit\n",
|
||||
"df[['STUDY', 'PHASE', 'LEVEL']].value_counts(ascending=True)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "f83b595c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"high_nback = df[\n",
|
||||
" (df[\"STUDY\"]==\"n-back\") &\n",
|
||||
" (df[\"LEVEL\"].isin([2, 3, 5, 6])) &\n",
|
||||
" (df[\"PHASE\"].isin([\"train\", \"test\"]))\n",
|
||||
"]\n",
|
||||
"high_nback.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "c0940343",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"low_all = df[\n",
|
||||
" ((df[\"PHASE\"] == \"baseline\") |\n",
|
||||
" ((df[\"STUDY\"] == \"n-back\") & (df[\"PHASE\"] != \"baseline\") & (df[\"LEVEL\"].isin([1,4]))))\n",
|
||||
"]\n",
|
||||
"print(low_all.shape)\n",
|
||||
"high_kdrive = df[\n",
|
||||
" (df[\"STUDY\"] == \"k-drive\") & (df[\"PHASE\"] != \"baseline\")\n",
|
||||
"]\n",
|
||||
"print(high_kdrive.shape)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "f7ce38d3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"print((df.shape[0]==(high_kdrive.shape[0]+high_nback.shape[0]+low_all.shape[0])))\n",
|
||||
"print(df.shape[0])\n",
|
||||
"print((high_kdrive.shape[0]+high_nback.shape[0]+low_all.shape[0]))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "48ba0379",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"high_all = pd.concat([high_nback, high_kdrive])\n",
|
||||
"high_all.shape"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "77dda26c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"print(f\"Gesamt: {df.shape[0]}=={low_all.shape[0]+high_all.shape[0]}\")\n",
|
||||
"print(f\"Anzahl an low load Samples: {low_all.shape[0]}\")\n",
|
||||
"print(f\"Anzahl an high load Samples: {high_all.shape[0]}\")\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "base",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.11.5"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -1,8 +1,10 @@
|
||||
import os
|
||||
import pandas as pd
|
||||
from pathlib import Path
|
||||
# TODO: Set paths correctly
|
||||
data_dir = Path("") # path to the directory with all .h5 files
|
||||
base_dir = Path(r"") # directory to store the parquet files in
|
||||
|
||||
data_dir = Path("/home/jovyan/Fahrsimulator_MSY2526_AI/EDA")
|
||||
|
||||
# Get all .h5 files and sort them
|
||||
matching_files = sorted(data_dir.glob("*.h5"))
|
||||
@@ -11,8 +13,8 @@ matching_files = sorted(data_dir.glob("*.h5"))
|
||||
CHUNK_SIZE = 50_000
|
||||
|
||||
for i, file_path in enumerate(matching_files):
|
||||
print(f"Subject {i} gestartet")
|
||||
print(f"{file_path} geoeffnet")
|
||||
print(f"Starting with subject {i}")
|
||||
print(f"Opened: {file_path}")
|
||||
|
||||
# Step 1: Get total number of rows and column names
|
||||
with pd.HDFStore(file_path, mode="r") as store:
|
||||
@@ -56,16 +58,16 @@ for i, file_path in enumerate(matching_files):
|
||||
start=start_row,
|
||||
stop=stop_row
|
||||
)
|
||||
|
||||
# print(f"[DEBUG] Vor Dropna: {df_chunk["EYE_LEFT_PUPIL_VALIDITY"].value_counts()}")
|
||||
# Add metadata columns
|
||||
df_chunk["subjectID"] = i
|
||||
df_chunk["rowID"] = range(start_row, stop_row)
|
||||
|
||||
# Clean data
|
||||
df_chunk = df_chunk[df_chunk["LEVEL"] != 0]
|
||||
df_chunk = df_chunk.dropna()
|
||||
# problematisch, weil die eye tracking auflösung kaputt geht
|
||||
df_chunk = df_chunk.dropna(subset=face_au_cols)
|
||||
|
||||
# print(f"[DEBUG] Nach Dropna: {df_chunk["EYE_LEFT_PUPIL_VALIDITY"].value_counts()}")
|
||||
# Only keep non-empty chunks
|
||||
if len(df_chunk) > 0:
|
||||
chunks_to_save.append(df_chunk)
|
||||
@@ -81,7 +83,7 @@ for i, file_path in enumerate(matching_files):
|
||||
print(f"Final dataframe shape: {df_final.shape}")
|
||||
|
||||
# Save to parquet
|
||||
base_dir = Path(r"/home/jovyan/data-paulusjafahrsimulator-gpu/both_mod_parquet_files")
|
||||
|
||||
os.makedirs(base_dir, exist_ok=True)
|
||||
|
||||
out_name = base_dir / f"both_mod_{i:04d}.parquet"
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -0,0 +1,528 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "47f6de7b",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Bibliotheken importieren"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "99294260",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd \n",
|
||||
"import numpy as np \n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import seaborn as sns \n",
|
||||
"import random \n",
|
||||
"import joblib \n",
|
||||
"from pathlib import Path \n",
|
||||
"\n",
|
||||
"from sklearn.model_selection import GroupKFold, GroupShuffleSplit\n",
|
||||
"from sklearn.preprocessing import StandardScaler \n",
|
||||
"from sklearn.metrics import ( \n",
|
||||
" precision_score, recall_score,\n",
|
||||
" confusion_matrix, roc_curve, auc, \n",
|
||||
" precision_recall_curve, f1_score, \n",
|
||||
" balanced_accuracy_score, accuracy_score\n",
|
||||
") \n",
|
||||
"\n",
|
||||
"import tensorflow as tf \n",
|
||||
"from tensorflow.keras import Input, layers, models, regularizers"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "52b4ca8c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Seed festlegen"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "6e49d281",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"SEED = 42 \n",
|
||||
"np.random.seed(SEED) \n",
|
||||
"tf.random.set_seed(SEED) \n",
|
||||
"random.seed(SEED)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "ae1a715f",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Daten laden"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "870f01c3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"data_path = Path(r\"~/data-paulusjafahrsimulator-gpu/new_datasets/50s_25Hz_dataset.parquet\") \n",
|
||||
"\n",
|
||||
"data = pd.read_parquet(path=data_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "bedbc23b",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Labels erstellen"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "38848515",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"low_all = data[((data[\"PHASE\"] == \"baseline\") | \n",
|
||||
" ((data[\"STUDY\"] == \"n-back\") & (data[\"PHASE\"] != \"baseline\") & (data[\"LEVEL\"].isin([1,4]))))].copy() \n",
|
||||
"\n",
|
||||
"high_all = pd.concat([ \n",
|
||||
" data[(data[\"STUDY\"]==\"n-back\") & (data[\"LEVEL\"].isin([2,3,5,6])) & (data[\"PHASE\"].isin([\"train\",\"test\"]))], \n",
|
||||
" data[(data[\"STUDY\"]==\"k-drive\") & (data[\"PHASE\"]!=\"baseline\")] \n",
|
||||
"]).copy() \n",
|
||||
"\n",
|
||||
"low_all[\"label\"] = 0 \n",
|
||||
"high_all[\"label\"] = 1 \n",
|
||||
"data = pd.concat([low_all, high_all], ignore_index=True).drop_duplicates() "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "0b282acf",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Features und Labels"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "5edb00a0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"#Face AUs\n",
|
||||
"au_columns = [col for col in data.columns if \"face\" in col.lower()] \n",
|
||||
"\n",
|
||||
"#Eye Features\n",
|
||||
"eye_columns = [ \n",
|
||||
" 'Fix_count_short_66_150', \n",
|
||||
" 'Fix_count_medium_300_500', \n",
|
||||
" 'Fix_count_long_gt_1000', \n",
|
||||
" 'Fix_count_100', \n",
|
||||
" 'Fix_mean_duration', \n",
|
||||
" 'Fix_median_duration', \n",
|
||||
" 'Sac_count', \n",
|
||||
" 'Sac_mean_amp', \n",
|
||||
" 'Sac_mean_dur', \n",
|
||||
" 'Sac_median_dur', \n",
|
||||
" 'Blink_count', \n",
|
||||
" 'Blink_mean_dur', \n",
|
||||
" 'Blink_median_dur', \n",
|
||||
" 'Pupil_mean', \n",
|
||||
" 'Pupil_IPA' \n",
|
||||
"]\n",
|
||||
"\n",
|
||||
"#Early Fusion\n",
|
||||
"feature_columns = au_columns + eye_columns\n",
|
||||
"\n",
|
||||
"#NaNs entfernen \n",
|
||||
"data = data.dropna(subset=feature_columns + [\"label\"])\n",
|
||||
"\n",
|
||||
"X = data[feature_columns].values[..., np.newaxis] \n",
|
||||
"y = data[\"label\"].values \n",
|
||||
"\n",
|
||||
"groups = data[\"subjectID\"].values\n",
|
||||
"print(data.columns.tolist())\n",
|
||||
"\n",
|
||||
"print(\"Gefundene FACE_AU-Spalten:\", au_columns)\n",
|
||||
"print(\"Gefundene Eye Features:\" , eye_columns)\n",
|
||||
"\n",
|
||||
"print(\"Anzahl FACE_AUs:\", len(au_columns)) \n",
|
||||
"print(\"Anzahl EYE Features:\", len(eye_columns)) \n",
|
||||
"print(\"Gesamtzahl Features:\", len(feature_columns))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "d8689679",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Train-Test-Split"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b5cf88c3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"gss = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42)\n",
|
||||
"train_idx, test_idx = next(gss.split(X, y, groups))\n",
|
||||
"\n",
|
||||
"#feature_columns_train, feature_columns_test = X[train_idx], X[test_idx]\n",
|
||||
"X_train, X_test = X[train_idx], X[test_idx]\n",
|
||||
"y_train, y_test = y[train_idx], y[test_idx]\n",
|
||||
"groups_train, groups_test = groups[train_idx], groups[test_idx]\n",
|
||||
"\n",
|
||||
"print(\"Train:\", len(y_train), \" | Test:\", len(y_test))\n",
|
||||
"print(\"Train:\", len(X_train), \" | Test:\", len(X_test))\n",
|
||||
"print(train_idx)\n",
|
||||
"print(test_idx)\n",
|
||||
"print(np.intersect1d(train_idx,test_idx))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "a539b83b",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"CNN-Modell"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e4a7f496",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def build_model(input_shape, lr=1e-4): \n",
|
||||
" model = models.Sequential([ \n",
|
||||
" Input(shape=input_shape), \n",
|
||||
" layers.Conv1D(32, kernel_size=3, activation=\"relu\", kernel_regularizer=regularizers.l2(0.001)), \n",
|
||||
" layers.BatchNormalization(), \n",
|
||||
" layers.MaxPooling1D(pool_size=2),\n",
|
||||
"\n",
|
||||
" layers.Conv1D(64, kernel_size=3, activation=\"relu\", kernel_regularizer=regularizers.l2(0.001)), \n",
|
||||
" layers.BatchNormalization(), \n",
|
||||
" layers.GlobalAveragePooling1D(), \n",
|
||||
" \n",
|
||||
" layers.Dense(32, activation=\"relu\", kernel_regularizer=regularizers.l2(0.001)), \n",
|
||||
" layers.Dropout(0.5), \n",
|
||||
" layers.Dense(1, activation=\"sigmoid\") \n",
|
||||
" ]) \n",
|
||||
" \n",
|
||||
" model.compile( \n",
|
||||
" optimizer=tf.keras.optimizers.Adam(learning_rate=lr), \n",
|
||||
" loss=\"binary_crossentropy\", \n",
|
||||
" metrics=[\"accuracy\", tf.keras.metrics.AUC(name=\"auc\")] \n",
|
||||
" ) \n",
|
||||
" return model"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "5905871b",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Cross-Validation"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "90658000",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"gkf = GroupKFold(n_splits=5) \n",
|
||||
"cv_histories = [] \n",
|
||||
"cv_results = [] \n",
|
||||
"fold_subjects = []\n",
|
||||
"all_conf_matrices = []\n",
|
||||
"\n",
|
||||
"for fold, (tr_idx, val_idx) in enumerate(gkf.split(X_train, y_train, groups_train)):\n",
|
||||
" train_subjects = np.unique(groups_train[tr_idx]) \n",
|
||||
" val_subjects = np.unique(groups_train[val_idx]) \n",
|
||||
" fold_subjects.append({\"Fold\": fold+1, \n",
|
||||
" \"Train_Subjects\": train_subjects, \n",
|
||||
" \"Val_Subjects\": val_subjects}) \n",
|
||||
" \n",
|
||||
" print(f\"\\n--- Fold {fold+1} ---\") \n",
|
||||
" print(\"Train-Subjects:\", train_subjects) \n",
|
||||
" print(\"Val-Subjects:\", val_subjects) \n",
|
||||
"\n",
|
||||
" #Split\n",
|
||||
" X_tr, X_val = X_train[tr_idx], X_train[val_idx] \n",
|
||||
" y_tr, y_val = y_train[tr_idx], y_train[val_idx] # Normalisierung pro Fold \n",
|
||||
"\n",
|
||||
" #Normalisierung pro Fold\n",
|
||||
" scaler = StandardScaler() \n",
|
||||
" X_tr = scaler.fit_transform(X_tr.reshape(len(X_tr), -1)).reshape(X_tr.shape) \n",
|
||||
" X_val = scaler.transform(X_val.reshape(len(X_val), -1)).reshape(X_val.shape) \n",
|
||||
"\n",
|
||||
" # Modell \n",
|
||||
" model = build_model(input_shape=(len(feature_columns),1), lr=1e-4) \n",
|
||||
" #model.summary() \n",
|
||||
"\n",
|
||||
" callbacks = [ \n",
|
||||
" tf.keras.callbacks.EarlyStopping(monitor=\"val_loss\", patience=10, restore_best_weights=True), \n",
|
||||
" tf.keras.callbacks.ReduceLROnPlateau(monitor=\"val_loss\", factor=0.5, patience=5, min_lr=1e-6) \n",
|
||||
" ] \n",
|
||||
"\n",
|
||||
" history = model.fit( \n",
|
||||
" X_tr, y_tr, \n",
|
||||
" validation_data=(X_val, y_val), \n",
|
||||
" epochs=50,\n",
|
||||
" batch_size=16, \n",
|
||||
" callbacks=callbacks, \n",
|
||||
" verbose=0 \n",
|
||||
" ) \n",
|
||||
"\n",
|
||||
" cv_histories.append(history.history) \n",
|
||||
" scores = model.evaluate(X_val, y_val, verbose=0) \n",
|
||||
" cv_results.append(scores) \n",
|
||||
" print(f\"Fold {fold+1} - Val Loss: {scores[0]:.4f}, Val Acc: {scores[1]:.4f}, Val AUC: {scores[2]:.4f}\")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" #Konfusionsmatrix \n",
|
||||
" y_pred = (model.predict(X_val) > 0.5).astype(int) \n",
|
||||
" cm = confusion_matrix(y_val, y_pred) \n",
|
||||
" all_conf_matrices.append(cm) \n",
|
||||
" \n",
|
||||
" print(f\"Konfusionsmatrix Fold {fold+1}:\\n{cm}\\n\") \n",
|
||||
" \n",
|
||||
"# Aggregierte Matrix \n",
|
||||
"agg_cm = sum(all_conf_matrices) \n",
|
||||
"print(\"Aggregierte Konfusionsmatrix über alle Folds:\") \n",
|
||||
"print(agg_cm)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "d10b7e78",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Results"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9aeba7f4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"#results\n",
|
||||
"cv_results = np.array(cv_results) \n",
|
||||
"print(\"\\n=== Cross-Validation Ergebnisse ===\") \n",
|
||||
"print(f\"Durchschnittlicher Val-Loss: {cv_results[:,0].mean():.4f}\") \n",
|
||||
"print(f\"Durchschnittliche Val-Accuracy: {cv_results[:,1].mean():.4f}\") \n",
|
||||
"print(f\"Durchschnittliche Val-AUC: {cv_results[:,2].mean():.4f}\")\n",
|
||||
"\n",
|
||||
"#Ergebnis-Tabelle erstellen\n",
|
||||
"results_table = pd.DataFrame({ \n",
|
||||
" \"Fold\": np.arange(1, len(cv_results)+1), \n",
|
||||
" \"Val Loss\": cv_results[:,0], \n",
|
||||
" \"Val Accuracy\": cv_results[:,1], \n",
|
||||
" \"Val AUC\": cv_results[:,2] }) \n",
|
||||
"\n",
|
||||
"# Durchschnittszeile hinzufügen \n",
|
||||
"avg_row = pd.DataFrame({ \n",
|
||||
" \"Fold\": [\"Ø\"], \n",
|
||||
" \"Val Loss\": [cv_results[:,0].mean()], \n",
|
||||
" \"Val Accuracy\": [cv_results[:,1].mean()], \n",
|
||||
" \"Val AUC\": [cv_results[:,2].mean()] \n",
|
||||
"}) \n",
|
||||
"\n",
|
||||
"results_table = pd.concat([results_table, avg_row], ignore_index=True) \n",
|
||||
"\n",
|
||||
"print(\"\\n=== Ergebnis-Tabelle ===\") \n",
|
||||
"print(results_table) \n",
|
||||
"\n",
|
||||
"#Tabelle speichern \n",
|
||||
"results_table.to_csv(\"cnn_crossVal_results.csv\", index=False) \n",
|
||||
"print(\"Ergebnisse gespeichert als 'cnn_crossVal_results.csv'\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "fae5df7a",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Finales Modell trainieren"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "5b3eab61",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"scaler_final = StandardScaler() \n",
|
||||
"X_train_scaled = scaler_final.fit_transform( X_train.reshape(len(X_train), -1) ).reshape(X_train.shape)\n",
|
||||
"\n",
|
||||
"final_model = build_model(input_shape=(len(feature_columns),1), lr=1e-4) \n",
|
||||
"#final_model.summary() \n",
|
||||
"\n",
|
||||
"final_model.fit( \n",
|
||||
" X_train_scaled, y_train,\n",
|
||||
" epochs=50, \n",
|
||||
" batch_size=16, \n",
|
||||
" verbose=1 \n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "7c7f9cc4",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Speichern des Modells"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "2d3af5be",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"final_model.save(\"cnn_crossVal_EarlyFusion_V2_0103.keras\") \n",
|
||||
"joblib.dump(scaler_final, \"scaler_crossVal_EarlyFusion_V2_0103.joblib\") \n",
|
||||
"\n",
|
||||
"# print(\"Finales Modell und Scaler gespeichert als 'cnn_crossVal_EarlyFusion_V2.keras' und 'scaler_crossVal_EarlyFusion_V2.joblib'\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "c11891e0",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Plots"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9f6a8584",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"#plots\n",
|
||||
"def plot_cv_histories(cv_histories, metric): \n",
|
||||
" plt.figure(figsize=(10,6)) \n",
|
||||
" \n",
|
||||
" for i, hist in enumerate(cv_histories): \n",
|
||||
" plt.plot(hist[metric], label=f\"Fold {i+1} Train\", alpha=0.7) \n",
|
||||
" plt.plot(hist[f\"val_{metric}\"], label=f\"Fold {i+1} Val\", linestyle=\"--\", alpha=0.7) \n",
|
||||
" plt.xlabel(\"Epochs\") \n",
|
||||
" plt.ylabel(metric.capitalize()) \n",
|
||||
" plt.title(f\"Cross-Validation {metric.capitalize()} Verläufe\") \n",
|
||||
" plt.legend() \n",
|
||||
" plt.grid(True) \n",
|
||||
" plt.show()\n",
|
||||
" \n",
|
||||
"plot_cv_histories(cv_histories, \"loss\") \n",
|
||||
"plot_cv_histories(cv_histories, \"accuracy\") \n",
|
||||
"plot_cv_histories(cv_histories, \"auc\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "4aebe6c6",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Test"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "0d34d6b7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Preprocessing Testdaten \n",
|
||||
"X_test_scaled = scaler.transform( \n",
|
||||
" X_test.reshape(len(X_test), -1) \n",
|
||||
").reshape(X_test.shape) \n",
|
||||
"\n",
|
||||
"# Vorhersagen \n",
|
||||
"y_prob_test = model.predict(X_test_scaled).flatten() \n",
|
||||
"y_pred_test = (y_prob_test > 0.5).astype(int) \n",
|
||||
"\n",
|
||||
"# Konfusionsmatrix \n",
|
||||
"cm_test = confusion_matrix(y_test, y_pred_test) \n",
|
||||
"\n",
|
||||
"plt.figure(figsize=(6,5)) \n",
|
||||
"sns.heatmap(cm_test, annot=True, fmt=\"d\", cmap=\"Greens\", \n",
|
||||
" xticklabels=[\"Pred 0\", \"Pred 1\"], \n",
|
||||
" yticklabels=[\"True 0\", \"True 1\"]) \n",
|
||||
"plt.title(\"Konfusionsmatrix - Testdaten\") \n",
|
||||
"plt.show() \n",
|
||||
"\n",
|
||||
"# ROC \n",
|
||||
"fpr, tpr, _ = roc_curve(y_test, y_prob_test) \n",
|
||||
"roc_auc = auc(fpr, tpr) \n",
|
||||
"\n",
|
||||
"plt.figure(figsize=(7,6)) \n",
|
||||
"plt.plot(fpr, tpr, label=f\"AUC = {roc_auc:.3f}\") \n",
|
||||
"plt.plot([0,1], [0,1], \"k--\") \n",
|
||||
"plt.title(\"ROC - Testdaten\") \n",
|
||||
"plt.legend() \n",
|
||||
"plt.grid(True) \n",
|
||||
"plt.show() \n",
|
||||
"\n",
|
||||
"# Precision-Recall \n",
|
||||
"precision, recall, _ = precision_recall_curve(y_test, y_prob_test) \n",
|
||||
"plt.figure(figsize=(7,6)) \n",
|
||||
"plt.plot(recall, precision) \n",
|
||||
"plt.title(\"Precision-Recall - Testdaten\") \n",
|
||||
"plt.grid(True) \n",
|
||||
"plt.show() \n",
|
||||
"\n",
|
||||
"# Metriken \n",
|
||||
"print(\"Accuracy:\", accuracy_score(y_test, y_pred_test))\n",
|
||||
"print(\"F1-Score:\", f1_score(y_test, y_pred_test)) \n",
|
||||
"print(\"Balanced Accuracy:\", balanced_accuracy_score(y_test, y_pred_test)) \n",
|
||||
"print(\"Precision:\", precision_score(y_test, y_pred_test)) \n",
|
||||
"print(\"Recall:\", recall_score(y_test, y_pred_test)) \n",
|
||||
"print(\"AUC:\", roc_auc)"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.12.10"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
File diff suppressed because one or more lines are too long
@@ -0,0 +1,472 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "b65b6b7d",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Imports"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "530e70af",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd \n",
|
||||
"import numpy as np \n",
|
||||
"import matplotlib.pyplot as plt \n",
|
||||
"import seaborn as sns \n",
|
||||
"import random \n",
|
||||
"import joblib \n",
|
||||
"from pathlib import Path \n",
|
||||
"\n",
|
||||
"from sklearn.model_selection import GroupKFold, GroupShuffleSplit \n",
|
||||
"from sklearn.preprocessing import StandardScaler \n",
|
||||
"from sklearn.metrics import ( \n",
|
||||
" precision_score, recall_score,\n",
|
||||
" confusion_matrix, roc_curve, auc, \n",
|
||||
" precision_recall_curve, f1_score, \n",
|
||||
" balanced_accuracy_score, accuracy_score\n",
|
||||
") \n",
|
||||
"\n",
|
||||
"import tensorflow as tf \n",
|
||||
"from tensorflow.keras import Input, layers, models"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "0d01127c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Seed"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "67aaf56e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"SEED = 42 \n",
|
||||
"np.random.seed(SEED) \n",
|
||||
"tf.random.set_seed(SEED) \n",
|
||||
"random.seed(SEED)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "844e250c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Daten laden "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "73a34b69",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"data_path = Path(r\"~/data-paulusjafahrsimulator-gpu/new_datasets/50s_25Hz_dataset.parquet\") \n",
|
||||
"data = pd.read_parquet(path=data_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "325179d3",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Daten vorbereiten"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "a5ad3126",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"low_all = data[ \n",
|
||||
" ((data[\"PHASE\"] == \"baseline\") | \n",
|
||||
" ((data[\"STUDY\"] == \"n-back\") & \n",
|
||||
" (data[\"PHASE\"] != \"baseline\") & \n",
|
||||
" (data[\"LEVEL\"].isin([1, 4])))) \n",
|
||||
"].copy() \n",
|
||||
"\n",
|
||||
"high_all = pd.concat([ \n",
|
||||
" data[(data[\"STUDY\"] == \"n-back\") & \n",
|
||||
" (data[\"LEVEL\"].isin([2, 3, 5, 6])) & \n",
|
||||
" (data[\"PHASE\"].isin([\"train\", \"test\"]))], \n",
|
||||
" data[(data[\"STUDY\"] == \"k-drive\") & (data[\"PHASE\"] != \"baseline\")] \n",
|
||||
"]).copy() \n",
|
||||
"\n",
|
||||
"low_all[\"label\"] = 0 \n",
|
||||
"high_all[\"label\"] = 1 \n",
|
||||
"\n",
|
||||
"data = pd.concat([low_all, high_all], ignore_index=True).drop_duplicates()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "fd843b62",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Features"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "5f10e6ca",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"au_columns = [col for col in data.columns if \"face\" in col.lower()] \n",
|
||||
"\n",
|
||||
"eye_columns = [ \n",
|
||||
" 'Fix_count_short_66_150','Fix_count_medium_300_500','Fix_count_long_gt_1000', \n",
|
||||
" 'Fix_count_100','Fix_mean_duration','Fix_median_duration', \n",
|
||||
" 'Sac_count','Sac_mean_amp','Sac_mean_dur','Sac_median_dur', \n",
|
||||
" 'Blink_count','Blink_mean_dur','Blink_median_dur', \n",
|
||||
" 'Pupil_mean','Pupil_IPA' \n",
|
||||
"] \n",
|
||||
"\n",
|
||||
"# NaNs entfernen \n",
|
||||
"data = data.dropna(subset=au_columns + eye_columns + [\"label\"]) \n",
|
||||
"\n",
|
||||
"# Arrays \n",
|
||||
"print(data[au_columns].shape)\n",
|
||||
"X_au = data[au_columns].values[..., np.newaxis] \n",
|
||||
"X_eye = data[eye_columns].values\n",
|
||||
"print(X_au.shape)\n",
|
||||
"y = data[\"label\"].values \n",
|
||||
"groups = data[\"subjectID\"].values"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "cabe09af",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Train/Test Split"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "52d3b7cf",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"gss = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42)\n",
|
||||
"train_idx, test_idx = next(gss.split(X_au, y, groups))\n",
|
||||
"\n",
|
||||
"X_au_train, X_au_test = X_au[train_idx], X_au[test_idx]\n",
|
||||
"X_eye_train, X_eye_test = X_eye[train_idx], X_eye[test_idx]\n",
|
||||
"y_train, y_test = y[train_idx], y[test_idx]\n",
|
||||
"groups_train, groups_test = groups[train_idx], groups[test_idx]\n",
|
||||
"\n",
|
||||
"print(\"Train:\", len(y_train), \" | Test:\", len(y_test))\n",
|
||||
"print(np.unique(groups_test))\n",
|
||||
"print(np.unique(groups_train))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6dedded5",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Hybrid CNN-Modell"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "41cc1b30",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def build_hybrid_model(n_aus, n_eye, lr=1e-4): \n",
|
||||
" input_au = Input(shape=(n_aus, 1), name=\"au_input\") \n",
|
||||
" x = layers.Conv1D(32, 3, activation=\"relu\")(input_au) \n",
|
||||
" x = layers.BatchNormalization()(x) \n",
|
||||
" x = layers.MaxPooling1D(2)(x) \n",
|
||||
" x = layers.Conv1D(64, 3, activation=\"relu\")(x) \n",
|
||||
" x = layers.BatchNormalization()(x) \n",
|
||||
" x = layers.GlobalAveragePooling1D()(x) \n",
|
||||
"\n",
|
||||
" input_eye = Input(shape=(n_eye,), name=\"eye_input\") \n",
|
||||
" e = layers.Dense(32, activation=\"relu\")(input_eye) \n",
|
||||
" e = layers.Dropout(0.3)(e) \n",
|
||||
" e = layers.Dense(16, activation=\"relu\")(e) \n",
|
||||
"\n",
|
||||
" fused = layers.concatenate([x, e]) \n",
|
||||
" z = layers.Dense(32, activation=\"relu\")(fused) \n",
|
||||
" z = layers.Dropout(0.4)(z) \n",
|
||||
" output = layers.Dense(1, activation=\"sigmoid\")(z) \n",
|
||||
"\n",
|
||||
" model = models.Model(inputs=[input_au, input_eye], outputs=output) \n",
|
||||
" model.compile( \n",
|
||||
" optimizer=tf.keras.optimizers.Adam(learning_rate=lr), \n",
|
||||
" loss=\"binary_crossentropy\", \n",
|
||||
" metrics=[\"accuracy\", tf.keras.metrics.AUC(name=\"auc\")] \n",
|
||||
" ) \n",
|
||||
" \n",
|
||||
" return model"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "cea6d0d0",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Cross Validation (nur Trainingsdaten)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9c390b46",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"gkf = GroupKFold(n_splits=5) \n",
|
||||
"cv_histories = [] \n",
|
||||
"cv_results = [] \n",
|
||||
"all_conf_matrices = [] \n",
|
||||
"fold_subjects = []\n",
|
||||
"\n",
|
||||
"for fold, (tr_idx, va_idx) in enumerate(gkf.split(X_au_train, y_train, groups_train)): \n",
|
||||
" \n",
|
||||
" train_subjects = np.unique(groups_train[tr_idx]) \n",
|
||||
" val_subjects = np.unique(groups_train[va_idx]) \n",
|
||||
" fold_subjects.append({\"Fold\": fold+1, \n",
|
||||
" \"Train_Subjects\": train_subjects, \n",
|
||||
" \"Val_Subjects\": val_subjects}) \n",
|
||||
" \n",
|
||||
" print(f\"\\n--- Fold {fold+1} ---\") \n",
|
||||
" print(\"Train-Subjects:\", train_subjects) \n",
|
||||
" print(\"Val-Subjects:\", val_subjects) \n",
|
||||
" \n",
|
||||
" X_tr_au, X_va_au = X_au_train[tr_idx], X_au_train[va_idx] \n",
|
||||
" X_tr_eye, X_va_eye = X_eye_train[tr_idx], X_eye_train[va_idx] \n",
|
||||
" y_tr, y_va = y_train[tr_idx], y_train[va_idx] \n",
|
||||
" \n",
|
||||
" # Scaler pro Fold \n",
|
||||
" scaler_au = StandardScaler() \n",
|
||||
" scaler_eye = StandardScaler() \n",
|
||||
" \n",
|
||||
" X_tr_au = scaler_au.fit_transform(X_tr_au.reshape(len(X_tr_au), -1)).reshape(X_tr_au.shape) \n",
|
||||
" X_va_au = scaler_au.transform(X_va_au.reshape(len(X_va_au), -1)).reshape(X_va_au.shape) \n",
|
||||
" \n",
|
||||
" X_tr_eye = scaler_eye.fit_transform(X_tr_eye) \n",
|
||||
" X_va_eye = scaler_eye.transform(X_va_eye) \n",
|
||||
" \n",
|
||||
" # Modell \n",
|
||||
" model_cv = build_hybrid_model(len(au_columns), len(eye_columns)) \n",
|
||||
" \n",
|
||||
" callbacks = [ \n",
|
||||
" tf.keras.callbacks.EarlyStopping(monitor=\"val_loss\", patience=10, restore_best_weights=True), \n",
|
||||
" tf.keras.callbacks.ReduceLROnPlateau(monitor=\"val_loss\", factor=0.5, patience=5, min_lr=1e-6) \n",
|
||||
" ] \n",
|
||||
" \n",
|
||||
" history = model_cv.fit( \n",
|
||||
" [X_tr_au, X_tr_eye], y_tr, \n",
|
||||
" validation_data=([X_va_au, X_va_eye], y_va), \n",
|
||||
" epochs=100, \n",
|
||||
" batch_size=16, \n",
|
||||
" verbose=0 \n",
|
||||
" ) \n",
|
||||
" \n",
|
||||
" cv_histories.append(history.history) \n",
|
||||
" \n",
|
||||
" # Evaluation \n",
|
||||
" scores = model_cv.evaluate([X_va_au, X_va_eye], y_va, verbose=0) \n",
|
||||
" cv_results.append(scores) \n",
|
||||
" print(f\"Val Loss={scores[0]:.4f} | Val Acc={scores[1]:.4f} | Val AUC={scores[2]:.4f}\") \n",
|
||||
" \n",
|
||||
" # Konfusionsmatrix pro Fold \n",
|
||||
" y_pred_va = (model_cv.predict([X_va_au, X_va_eye]) > 0.5).astype(int) \n",
|
||||
" cm = confusion_matrix(y_va, y_pred_va) \n",
|
||||
" all_conf_matrices.append(cm) \n",
|
||||
" \n",
|
||||
" plt.figure(figsize=(6,5)) \n",
|
||||
" sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", \n",
|
||||
" xticklabels=[\"Pred 0\", \"Pred 1\"], \n",
|
||||
" yticklabels=[\"True 0\", \"True 1\"]) \n",
|
||||
" plt.title(f\"Konfusionsmatrix - Fold {fold+1}\") \n",
|
||||
" plt.show() \n",
|
||||
" \n",
|
||||
"# Aggregierte Konfusionsmatrix \n",
|
||||
"agg_cm = sum(all_conf_matrices) \n",
|
||||
"\n",
|
||||
"plt.figure(figsize=(6,5)) \n",
|
||||
"sns.heatmap(agg_cm, annot=True, fmt=\"d\", cmap=\"Purples\", \n",
|
||||
" xticklabels=[\"Pred 0\", \"Pred 1\"], \n",
|
||||
" yticklabels=[\"True 0\", \"True 1\"]) \n",
|
||||
"plt.title(\"Aggregierte Konfusionsmatrix - alle Folds\") \n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "97df9df1",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Results"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9eae5c0f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"#results\n",
|
||||
"cv_results = np.array(cv_results) \n",
|
||||
"print(\"\\n=== Cross-Validation Ergebnisse ===\") \n",
|
||||
"print(f\"Durchschnittlicher Val-Loss: {cv_results[:,0].mean():.4f}\") \n",
|
||||
"print(f\"Durchschnittliche Val-Accuracy: {cv_results[:,1].mean():.4f}\") \n",
|
||||
"print(f\"Durchschnittliche Val-AUC: {cv_results[:,2].mean():.4f}\")\n",
|
||||
"\n",
|
||||
"#Ergebnis-Tabelle erstellen\n",
|
||||
"results_table = pd.DataFrame({ \n",
|
||||
" \"Fold\": np.arange(1, len(cv_results)+1), \n",
|
||||
" \"Val Loss\": cv_results[:,0], \n",
|
||||
" \"Val Accuracy\": cv_results[:,1], \n",
|
||||
" \"Val AUC\": cv_results[:,2] }) \n",
|
||||
"\n",
|
||||
"# Durchschnittszeile hinzufügen \n",
|
||||
"avg_row = pd.DataFrame({ \n",
|
||||
" \"Fold\": [\"Ø\"], \n",
|
||||
" \"Val Loss\": [cv_results[:,0].mean()], \n",
|
||||
" \"Val Accuracy\": [cv_results[:,1].mean()], \n",
|
||||
" \"Val AUC\": [cv_results[:,2].mean()] \n",
|
||||
"}) \n",
|
||||
"\n",
|
||||
"results_table = pd.concat([results_table, avg_row], ignore_index=True) \n",
|
||||
"\n",
|
||||
"print(\"\\n=== Ergebnis-Tabelle ===\") \n",
|
||||
"print(results_table) \n",
|
||||
"\n",
|
||||
"#Tabelle speichern \n",
|
||||
"results_table.to_csv(\"cnn_crossVal_results.csv\", index=False) \n",
|
||||
"print(\"Ergebnisse gespeichert als 'cnn_crossVal_results.csv'\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "7e564308",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Speichern des Modells"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9afc926b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"model_cv.save(\"hybrid_fusion_model_Test_group_split_0103.keras\") \n",
|
||||
"joblib.dump(scaler_au, \"scaler_au_Test_group_split_0103.joblib\") \n",
|
||||
"joblib.dump(scaler_eye, \"scaler_eye_Test_group_split_0103.joblib\") \n",
|
||||
"\n",
|
||||
"print(\"Finales Modell gespeichert.\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "391af5d5",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Test"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "0bb8c14c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Preprocessing Testdaten \n",
|
||||
"X_au_test_scaled = scaler_au.transform( \n",
|
||||
" X_au_test.reshape(len(X_au_test), -1) \n",
|
||||
").reshape(X_au_test.shape) \n",
|
||||
"\n",
|
||||
"X_eye_test_scaled = scaler_eye.transform(X_eye_test) \n",
|
||||
"\n",
|
||||
"# Vorhersagen \n",
|
||||
"y_prob_test = model_cv.predict([X_au_test_scaled, X_eye_test_scaled]).flatten() \n",
|
||||
"y_pred_test = (y_prob_test > 0.5).astype(int) \n",
|
||||
"\n",
|
||||
"# Konfusionsmatrix \n",
|
||||
"cm_test = confusion_matrix(y_test, y_pred_test) \n",
|
||||
"\n",
|
||||
"plt.figure(figsize=(6,5)) \n",
|
||||
"sns.heatmap(cm_test, annot=True, fmt=\"d\", cmap=\"Greens\", \n",
|
||||
" xticklabels=[\"Pred 0\", \"Pred 1\"], \n",
|
||||
" yticklabels=[\"True 0\", \"True 1\"]) \n",
|
||||
"plt.title(\"Konfusionsmatrix - Testdaten\") \n",
|
||||
"plt.show() \n",
|
||||
"\n",
|
||||
"# ROC \n",
|
||||
"fpr, tpr, _ = roc_curve(y_test, y_prob_test) \n",
|
||||
"roc_auc = auc(fpr, tpr) \n",
|
||||
"\n",
|
||||
"plt.figure(figsize=(7,6)) \n",
|
||||
"plt.plot(fpr, tpr, label=f\"AUC = {roc_auc:.3f}\") \n",
|
||||
"plt.plot([0,1], [0,1], \"k--\") \n",
|
||||
"plt.title(\"ROC - Testdaten\") \n",
|
||||
"plt.legend() \n",
|
||||
"plt.grid(True) \n",
|
||||
"plt.show() \n",
|
||||
"\n",
|
||||
"# Precision-Recall \n",
|
||||
"precision, recall, _ = precision_recall_curve(y_test, y_prob_test) \n",
|
||||
"plt.figure(figsize=(7,6)) \n",
|
||||
"plt.plot(recall, precision) \n",
|
||||
"plt.title(\"Precision-Recall - Testdaten\") \n",
|
||||
"plt.grid(True) \n",
|
||||
"plt.show() \n",
|
||||
"\n",
|
||||
"# Metriken \n",
|
||||
"print(\"Accuracy:\", accuracy_score(y_test, y_pred_test))\n",
|
||||
"print(\"F1-Score:\", f1_score(y_test, y_pred_test)) \n",
|
||||
"print(\"Balanced Accuracy:\", balanced_accuracy_score(y_test, y_pred_test)) \n",
|
||||
"print(\"Precision:\", precision_score(y_test, y_pred_test)) \n",
|
||||
"print(\"Recall:\", recall_score(y_test, y_pred_test)) \n",
|
||||
"print(\"AUC:\", roc_auc)"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.12.10"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
@@ -0,0 +1,308 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "d48f2e13",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Importe"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e34b838d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np \n",
|
||||
"import pandas as pd \n",
|
||||
"import joblib \n",
|
||||
"import seaborn as sns \n",
|
||||
"import matplotlib.pyplot as plt \n",
|
||||
"\n",
|
||||
"from sklearn.metrics import ( \n",
|
||||
" confusion_matrix, \n",
|
||||
" roc_curve, auc, \n",
|
||||
" precision_recall_curve, \n",
|
||||
" f1_score, \n",
|
||||
" balanced_accuracy_score \n",
|
||||
")\n",
|
||||
" \n",
|
||||
"import tensorflow as tf"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "324554b5",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Modell und Scaler laden"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4acc3d2f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"model = tf.keras.models.load_model(\"hybrid_fusion_model_V2.keras\") \n",
|
||||
"scaler_au = joblib.load(\"scaler_au_V2.joblib\") \n",
|
||||
"scaler_eye = joblib.load(\"scaler_eye_V2.joblib\")\n",
|
||||
"\n",
|
||||
"print(\"Modell & Scaler erfolgreich geladen.\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "4271cbee",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Features laden"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "8342ea10",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"au_columns = [...] \n",
|
||||
"eye_columns = [...]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "4a58b20c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Preprocessing"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b683be47",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def preprocess_sample(df, au_columns, eye_columns, scaler_au, scaler_eye):\n",
|
||||
" # AUs\n",
|
||||
" X_au = df[au_columns].values\n",
|
||||
" X_au = scaler_au.transform(X_au).reshape(len(df), len(au_columns), 1)\n",
|
||||
"\n",
|
||||
" # Eye\n",
|
||||
" X_eye = df[eye_columns].values\n",
|
||||
" X_eye = scaler_eye.transform(X_eye)\n",
|
||||
"\n",
|
||||
" return X_au, X_eye"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "9dc99a3d",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Predict-Funktion"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "00295aa6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def predict_workload(df, model, au_columns, eye_columns, scaler_au, scaler_eye):\n",
|
||||
" X_au, X_eye = preprocess_sample(df, au_columns, eye_columns, scaler_au, scaler_eye)\n",
|
||||
"\n",
|
||||
" probs = model.predict([X_au, X_eye]).flatten()\n",
|
||||
" preds = (probs > 0.5).astype(int)\n",
|
||||
" \n",
|
||||
" return preds, probs"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "5753516b",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Testdaten laden"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "8875b0ee",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"test_data = pd.read_csv(\"test_data.csv\") # oder direkt aus Notebook 1 exportieren \n",
|
||||
"\n",
|
||||
"X_au_test = test_data[au_columns].values[..., np.newaxis] \n",
|
||||
"X_eye_test = test_data[eye_columns].values \n",
|
||||
"y_test = test_data[\"label\"].values \n",
|
||||
"groups_test = test_data[\"subjectID\"].values \n",
|
||||
"\n",
|
||||
"X_au_test_scaled = scaler_au.transform(X_au_test.reshape(len(X_au_test), -1)).reshape(X_au_test.shape) \n",
|
||||
"X_eye_test_scaled = scaler_eye.transform(X_eye_test)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "332a3a07",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Vorhersagen"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b5f58ece",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"y_prob = model.predict([X_au_test_scaled, X_eye_test_scaled]).flatten() \n",
|
||||
"y_pred = (y_prob > 0.5).astype(int)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "3bc5c66c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Konfusionsmatrix"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "40648dd7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"cm = confusion_matrix(y_test, y_pred) \n",
|
||||
"plt.figure(figsize=(6,5)) \n",
|
||||
"sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", \n",
|
||||
" xticklabels=[\"Pred 0\", \"Pred 1\"], \n",
|
||||
" yticklabels=[\"True 0\", \"True 1\"]) \n",
|
||||
"plt.title(\"Konfusionsmatrix - Testdaten\") \n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "e79ad8a6",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"ROC"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "dd93f15c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"fpr, tpr, _ = roc_curve(y_test, y_prob) \n",
|
||||
"roc_auc = auc(fpr, tpr) \n",
|
||||
"\n",
|
||||
"plt.figure(figsize=(7,6)) \n",
|
||||
"plt.plot(fpr, tpr, label=f\"AUC = {roc_auc:.3f}\") \n",
|
||||
"plt.plot([0,1], [0,1], \"k--\") \n",
|
||||
"plt.xlabel(\"False Positive Rate\") \n",
|
||||
"plt.ylabel(\"True Positive Rate\") \n",
|
||||
"plt.title(\"ROC‑Kurve – Testdaten\") \n",
|
||||
"plt.legend() \n",
|
||||
"plt.grid(True) \n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "2eaaf2a0",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Precision-Recall"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "601e5dc9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"precision, recall, _ = precision_recall_curve(y_test, y_prob) \n",
|
||||
"plt.figure(figsize=(7,6)) \n",
|
||||
"plt.plot(recall, precision) \n",
|
||||
"plt.xlabel(\"Recall\") \n",
|
||||
"plt.ylabel(\"Precision\") \n",
|
||||
"plt.title(\"Precision‑Recall‑Kurve – Testdaten\")\n",
|
||||
"plt.grid(True) \n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "270af771",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Scores"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e2e7da5b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"print(\"F1‑Score:\", f1_score(y_test, y_pred)) \n",
|
||||
"print(\"Balanced Accuracy:\", balanced_accuracy_score(y_test, y_pred))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "c6e22e1a",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Subject-Performance"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "731aaf73",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_eval = pd.DataFrame({ \n",
|
||||
" \"subject\": groups_test, \n",
|
||||
" \"y_true\": y_test, \n",
|
||||
" \"y_pred\": y_pred \n",
|
||||
"}) \n",
|
||||
"\n",
|
||||
"subject_perf = df_eval.groupby(\"subject\").apply( \n",
|
||||
" lambda x: balanced_accuracy_score(x[\"y_true\"], x[\"y_pred\"]) \n",
|
||||
") \n",
|
||||
"\n",
|
||||
"print(\"\\n=== Balanced Accuracy pro Proband ===\") \n",
|
||||
"print(subject_perf.sort_values())"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -23,7 +23,7 @@
|
||||
"id": "bef91203",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Imports"
|
||||
"### Imports + GPU "
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -38,19 +38,49 @@
|
||||
"from pathlib import Path\n",
|
||||
"import sys\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"import time\n",
|
||||
"base_dir = os.path.abspath(os.path.join(os.getcwd(), \"..\"))\n",
|
||||
"sys.path.append(base_dir)\n",
|
||||
"print(base_dir)\n",
|
||||
"\n",
|
||||
"from Fahrsimulator_MSY2526_AI.model_training.tools import evaluation_tools, scaler, mad_outlier_removal\n",
|
||||
"from Fahrsimulator_MSY2526_AI.model_training.tools import evaluation_tools, scaler, mad_outlier_removal, performance_split\n",
|
||||
"from sklearn.preprocessing import StandardScaler, MinMaxScaler\n",
|
||||
"from sklearn.svm import OneClassSVM\n",
|
||||
"from sklearn.model_selection import GridSearchCV, KFold, ParameterGrid, train_test_split, GroupKFold\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import tensorflow as tf\n",
|
||||
"from tensorflow.keras import layers, models, regularizers\n",
|
||||
"import pickle\n",
|
||||
"from sklearn.metrics import (roc_auc_score, accuracy_score, precision_score, recall_score, f1_score, confusion_matrix, classification_report, balanced_accuracy_score, ConfusionMatrixDisplay) "
|
||||
"from sklearn.metrics import (accuracy_score, auc, roc_curve, f1_score) "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "f03c8da9",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Check GPU availability\n",
|
||||
"print(\"TensorFlow version:\", tf.__version__)\n",
|
||||
"print(\"GPU Available:\", tf.config.list_physical_devices('GPU'))\n",
|
||||
"print(\"CUDA Available:\", tf.test.is_built_with_cuda())\n",
|
||||
"\n",
|
||||
"# Get detailed GPU info\n",
|
||||
"gpus = tf.config.list_physical_devices('GPU')\n",
|
||||
"if gpus:\n",
|
||||
" print(f\"\\nNumber of GPUs: {len(gpus)}\")\n",
|
||||
" for gpu in gpus:\n",
|
||||
" print(f\"GPU: {gpu}\")\n",
|
||||
" \n",
|
||||
" # Enable memory growth to prevent TF from allocating all GPU memory\n",
|
||||
" try:\n",
|
||||
" for gpu in gpus:\n",
|
||||
" tf.config.experimental.set_memory_growth(gpu, True)\n",
|
||||
" print(\"\\nGPU memory growth enabled\")\n",
|
||||
" except RuntimeError as e:\n",
|
||||
" print(e)\n",
|
||||
"else:\n",
|
||||
" print(\"\\nNo GPU found - running on CPU\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -58,15 +88,40 @@
|
||||
"id": "f00a477c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Data Preprocessing"
|
||||
"### Configuration of paths and data preprocessing"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "504c1df7",
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "5136fcec",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"Laden der Daten"
|
||||
"# TODO: set path where to save normalizer\n",
|
||||
"normalizer_path=Path('.pkl') # TODO: set manually"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "c2115f65",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"performance_path = Path(r\".csv\") # TODO: set manually\n",
|
||||
"performance_df = pd.read_csv(performance_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "559eb8d2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"encoder_save_path = Path('.keras') # TODO: set manually\n",
|
||||
"deep_svdd_save_path = Path('.keras') # TODO: set manually"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -76,7 +131,7 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dataset_path = Path(r\"/home/jovyan/data-paulusjafahrsimulator-gpu/first_AU_dataset/output_windowed.parquet\")"
|
||||
"dataset_path = Path(r\".parquet\") # TODO: set manually"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -89,12 +144,218 @@
|
||||
"df = pd.read_parquet(path=dataset_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "c045c46d",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Performance based split"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "1660ec95",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_ids, temp_ids, diff1 = performance_split.performance_based_split(\n",
|
||||
" subject_ids=df[\"subjectID\"].unique(),\n",
|
||||
" performance_df=performance_df,\n",
|
||||
" split_ratio=0.6, # 60% train, 40% temp\n",
|
||||
" random_seed=42\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"val_ids, test_ids, diff2 = performance_split.performance_based_split(\n",
|
||||
" subject_ids=temp_ids,\n",
|
||||
" performance_df=performance_df,\n",
|
||||
" split_ratio=0.5, # 50/50 split of remaining 40%\n",
|
||||
" random_seed=43\n",
|
||||
")\n",
|
||||
"print(diff1, diff2)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "195b7283",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Labeling"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "05b6b73d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"low_all = df[\n",
|
||||
" ((df[\"PHASE\"] == \"baseline\") |\n",
|
||||
" ((df[\"STUDY\"] == \"n-back\") & (df[\"PHASE\"] != \"baseline\") & (df[\"LEVEL\"].isin([1, 4]))))\n",
|
||||
"]\n",
|
||||
"print(f\"low all: {low_all.shape}\")\n",
|
||||
"\n",
|
||||
"high_nback = df[\n",
|
||||
" (df[\"STUDY\"]==\"n-back\") &\n",
|
||||
" (df[\"LEVEL\"].isin([2, 3, 5, 6])) &\n",
|
||||
" (df[\"PHASE\"].isin([\"train\", \"test\"]))\n",
|
||||
"]\n",
|
||||
"print(f\"high n-back: {high_nback.shape}\")\n",
|
||||
"\n",
|
||||
"high_kdrive = df[\n",
|
||||
" (df[\"STUDY\"] == \"k-drive\") & (df[\"PHASE\"] != \"baseline\")\n",
|
||||
"]\n",
|
||||
"print(f\"high k-drive: {high_kdrive.shape}\")\n",
|
||||
"\n",
|
||||
"high_all = pd.concat([high_nback, high_kdrive])\n",
|
||||
"print(f\"high all: {high_all.shape}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "60148c0b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"low = low_all.copy()\n",
|
||||
"high = high_all.copy()\n",
|
||||
"\n",
|
||||
"low[\"label\"] = 0\n",
|
||||
"high[\"label\"] = 1\n",
|
||||
"\n",
|
||||
"data = pd.concat([low, high], ignore_index=True)\n",
|
||||
"df = data.drop_duplicates()\n",
|
||||
"df = df.dropna()\n",
|
||||
"print(\"Label distribution:\")\n",
|
||||
"print(df[\"label\"].value_counts())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "c8fefca7",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Split"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "da6a2f87",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_df = df[\n",
|
||||
" (df.subjectID.isin(train_ids)) & (df['label'] == 0)\n",
|
||||
"].copy()\n",
|
||||
"\n",
|
||||
"# Validation: balanced sampling of label=0 and label=1\n",
|
||||
"val_df_full = df[df.subjectID.isin(val_ids)].copy()\n",
|
||||
"\n",
|
||||
"# Get all label=0 samples\n",
|
||||
"val_df_label0 = val_df_full[val_df_full['label'] == 0]\n",
|
||||
"\n",
|
||||
"# Sample same number from label=1\n",
|
||||
"n_samples = len(val_df_label0)\n",
|
||||
"val_df_label1 = val_df_full[val_df_full['label'] == 1].sample(\n",
|
||||
" n=n_samples, random_state=42\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Combine\n",
|
||||
"val_df = pd.concat([val_df_label0, val_df_label1], ignore_index=True)\n",
|
||||
"test_df = df[df.subjectID.isin(test_ids)]\n",
|
||||
"print(train_df.shape, val_df.shape,test_df.shape)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e8375760",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"val_df['label'].value_counts()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "f0570a3c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Normalization"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "cdd2ba73",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"face_au_cols = [c for c in train_df.columns if c.startswith(\"FACE_AU\")]\n",
|
||||
"eye_cols = ['Fix_count_short_66_150', 'Fix_count_medium_300_500',\n",
|
||||
" 'Fix_count_long_gt_1000', 'Fix_count_100', 'Fix_mean_duration',\n",
|
||||
" 'Fix_median_duration', 'Sac_count', 'Sac_mean_amp', 'Sac_mean_dur',\n",
|
||||
" 'Sac_median_dur', 'Blink_count', 'Blink_mean_dur', 'Blink_median_dur',\n",
|
||||
" 'Pupil_mean', 'Pupil_IPA']\n",
|
||||
"print(len(eye_cols))\n",
|
||||
"all_signal_columns = face_au_cols+eye_cols\n",
|
||||
"print(len(all_signal_columns))\n",
|
||||
"\n",
|
||||
"# fit and save normalizer\n",
|
||||
"normalizer = scaler.fit_normalizer(train_df, all_signal_columns, method='minmax', scope='global')\n",
|
||||
"scaler.save_normalizer(normalizer, normalizer_path )"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "76afc4d3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"normalizer = scaler.load_normalizer(normalizer_path)\n",
|
||||
"# Apply normalization to all sets\n",
|
||||
"train_df_norm = scaler.apply_normalizer(train_df, all_signal_columns, normalizer)\n",
|
||||
"val_df_norm = scaler.apply_normalizer(val_df, all_signal_columns, normalizer)\n",
|
||||
"test_df_norm = scaler.apply_normalizer(test_df, all_signal_columns, normalizer)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "77deead9",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Outlier removal (later)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "fd139799",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Change of dtypes for keras pandas"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "8587343e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"X_face = train_df_norm[face_au_cols].to_numpy(dtype=np.float32)\n",
|
||||
"X_eye = train_df_norm[eye_cols].to_numpy(dtype=np.float32)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "b736bc58",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Modell Training"
|
||||
"### Autoencoder Pre-Training"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -104,6 +365,529 @@
|
||||
"source": [
|
||||
"Vor-Training der Gewichte mit Autoencoder, Loss: MSE"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "3eab9d94",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def build_intermediate_fusion_autoencoder(\n",
|
||||
" input_dim_mod1=15,\n",
|
||||
" input_dim_mod2=20,\n",
|
||||
" encoder_hidden_dim_mod1=12, # TODO: set manually\n",
|
||||
" encoder_hidden_dim_mod2=20, # TODO: set manually\n",
|
||||
" latent_dim=6, # TODO: set manually\n",
|
||||
" dropout_rate=0.4, # TODO: set manually\n",
|
||||
" neg_slope=0.1, # TODO: set manually\n",
|
||||
" weight_decay=1e-4, # TODO: set manually\n",
|
||||
" decoder_hidden_dims=[16, 32] # TODO: set manually\n",
|
||||
"):\n",
|
||||
" \"\"\"\n",
|
||||
" Verbesserter Intermediate-Fusion Autoencoder für Deep SVDD.\n",
|
||||
" Änderungen:\n",
|
||||
" - Bottleneck vergrößert (latent_dim)\n",
|
||||
" - Dropout nur in Hidden Layers, nicht im Bottleneck\n",
|
||||
" - Decoder größer für stabileres Pretraining\n",
|
||||
" - Parametrisierbare Hidden-Dimensions für Encoder\n",
|
||||
" \"\"\"\n",
|
||||
"\n",
|
||||
" l2 = regularizers.l2(weight_decay)\n",
|
||||
" act = layers.LeakyReLU(negative_slope=neg_slope)\n",
|
||||
"\n",
|
||||
" # -------- Inputs --------\n",
|
||||
" x1_in = layers.Input(shape=(input_dim_mod1,), name=\"modality_1\")\n",
|
||||
" x2_in = layers.Input(shape=(input_dim_mod2,), name=\"modality_2\")\n",
|
||||
"\n",
|
||||
" # -------- Encoder 1 --------\n",
|
||||
" e1 = layers.Dense(\n",
|
||||
" encoder_hidden_dim_mod1,\n",
|
||||
" use_bias=False,\n",
|
||||
" kernel_regularizer=l2\n",
|
||||
" )(x1_in)\n",
|
||||
" e1 = act(e1)\n",
|
||||
" e1 = layers.Dropout(dropout_rate)(e1) \n",
|
||||
"\n",
|
||||
" e1 = layers.Dense(\n",
|
||||
" 16, \n",
|
||||
" use_bias=False,\n",
|
||||
" kernel_regularizer=l2\n",
|
||||
" )(e1)\n",
|
||||
" e1 = act(e1)\n",
|
||||
"\n",
|
||||
" # -------- Encoder 2 --------\n",
|
||||
" e2 = layers.Dense(\n",
|
||||
" encoder_hidden_dim_mod2,\n",
|
||||
" use_bias=False,\n",
|
||||
" kernel_regularizer=l2\n",
|
||||
" )(x2_in)\n",
|
||||
" e2 = act(e2)\n",
|
||||
" e2 = layers.Dropout(dropout_rate)(e2) \n",
|
||||
"\n",
|
||||
" e2 = layers.Dense(\n",
|
||||
" 16, \n",
|
||||
" use_bias=False,\n",
|
||||
" kernel_regularizer=l2\n",
|
||||
" )(e2)\n",
|
||||
" e2 = act(e2)\n",
|
||||
"\n",
|
||||
" # -------- Intermediate Fusion --------\n",
|
||||
" fused = layers.Concatenate(name=\"fusion\")([e1, e2]) # 16+16=32 dimensions\n",
|
||||
"\n",
|
||||
" # -------- Joint Encoder / Bottleneck --------\n",
|
||||
"\n",
|
||||
" h = layers.Dense(\n",
|
||||
" latent_dim,\n",
|
||||
" use_bias=False,\n",
|
||||
" kernel_regularizer=l2\n",
|
||||
" )(fused)\n",
|
||||
" h = act(h)\n",
|
||||
" h = layers.Dropout(dropout_rate)(h)\n",
|
||||
"\n",
|
||||
" z = layers.Dense(\n",
|
||||
" latent_dim,\n",
|
||||
" activation=None, # linear for Deep SVDD\n",
|
||||
" use_bias=False,\n",
|
||||
" kernel_regularizer=l2,\n",
|
||||
" name=\"latent\"\n",
|
||||
" )(h)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
" # -------- Decoder --------\n",
|
||||
" d = layers.Dense(\n",
|
||||
" decoder_hidden_dims[0], \n",
|
||||
" use_bias=False,\n",
|
||||
" kernel_regularizer=l2\n",
|
||||
" )(z)\n",
|
||||
" d = act(d)\n",
|
||||
"\n",
|
||||
" d = layers.Dense(\n",
|
||||
" decoder_hidden_dims[1],\n",
|
||||
" use_bias=False,\n",
|
||||
" kernel_regularizer=l2\n",
|
||||
" )(d)\n",
|
||||
" d = act(d)\n",
|
||||
"\n",
|
||||
" x1_out = layers.Dense(\n",
|
||||
" input_dim_mod1,\n",
|
||||
" activation=None,\n",
|
||||
" use_bias=False,\n",
|
||||
" name=\"recon_modality_1\"\n",
|
||||
" )(d)\n",
|
||||
"\n",
|
||||
" x2_out = layers.Dense(\n",
|
||||
" input_dim_mod2,\n",
|
||||
" activation=None,\n",
|
||||
" use_bias=False,\n",
|
||||
" name=\"recon_modality_2\"\n",
|
||||
" )(d)\n",
|
||||
"\n",
|
||||
" model = models.Model(\n",
|
||||
" inputs=[x1_in, x2_in],\n",
|
||||
" outputs=[x1_out, x2_out],\n",
|
||||
" name=\"IntermediateFusionAE_Improved\"\n",
|
||||
" )\n",
|
||||
"\n",
|
||||
" return model\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "80cb8eb0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"model = build_intermediate_fusion_autoencoder(\n",
|
||||
" input_dim_mod1=len(face_au_cols),\n",
|
||||
" input_dim_mod2=len(eye_cols),\n",
|
||||
" encoder_hidden_dim_mod1=12, # TODO: set manually\n",
|
||||
" encoder_hidden_dim_mod2=8, # TODO: set manually\n",
|
||||
" latent_dim=4,\n",
|
||||
" dropout_rate=0.7, # TODO: set manually\n",
|
||||
" neg_slope=0.1,\n",
|
||||
" weight_decay=1e-3\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"model.compile(\n",
|
||||
" loss={\n",
|
||||
" \"recon_modality_1\": \"mse\",\n",
|
||||
" \"recon_modality_2\": \"mse\",\n",
|
||||
" },\n",
|
||||
" loss_weights={\n",
|
||||
" \"recon_modality_1\": 1.0,\n",
|
||||
" \"recon_modality_2\": 1.0,\n",
|
||||
" },\n",
|
||||
" optimizer=tf.keras.optimizers.Adam(1e-3)\n",
|
||||
" \n",
|
||||
")\n",
|
||||
"\n",
|
||||
"batch_size_ae=64\n",
|
||||
"# model.summary()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "95d36a07",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"model.fit(\n",
|
||||
" x=[X_face, X_eye],\n",
|
||||
" y=[X_face, X_eye],\n",
|
||||
" batch_size=batch_size_ae,\n",
|
||||
" epochs=150,\n",
|
||||
" shuffle=True\n",
|
||||
")\n",
|
||||
"model.compile(\n",
|
||||
" loss={\n",
|
||||
" \"recon_modality_1\": \"mse\",\n",
|
||||
" \"recon_modality_2\": \"mse\",\n",
|
||||
" },\n",
|
||||
" loss_weights={\n",
|
||||
" \"recon_modality_1\": 1.0,\n",
|
||||
" \"recon_modality_2\": 1.0,\n",
|
||||
" },\n",
|
||||
" optimizer=tf.keras.optimizers.Adam(1e-4),\n",
|
||||
")\n",
|
||||
"model.fit(\n",
|
||||
" x=[X_face, X_eye],\n",
|
||||
" y=[X_face, X_eye],\n",
|
||||
" batch_size=batch_size_ae,\n",
|
||||
" epochs=100,\n",
|
||||
" shuffle=True\n",
|
||||
")\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9ccfbc71",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"encoder = tf.keras.Model(\n",
|
||||
" inputs=model.inputs,\n",
|
||||
" outputs=model.get_layer(\"latent\").output,\n",
|
||||
" name=\"SVDD_Encoder\"\n",
|
||||
")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "e4e1b5ff",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Speichern"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "7e591264",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"encoder.save(encoder_save_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "372dc754",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Laden Encoder / Deepsvdd"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "83199fc6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"encoder_load_path = encoder_save_path\n",
|
||||
"encoder = tf.keras.models.load_model(encoder_load_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "92046112",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Check, if encoder works"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "db2fa21c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"ans= encoder.predict([X_face, X_eye])\n",
|
||||
"print(ans[:6,:])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "d7bcc35d",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Deep SVDD Training"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "806a2479",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"encoder_load_path = encoder_save_path\n",
|
||||
"deep_svdd_net = tf.keras.models.load_model(encoder_load_path) "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "54083759",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def get_center(model, dataset):\n",
|
||||
" center = model.predict(dataset).mean(axis=0)\n",
|
||||
"\n",
|
||||
" eps = 0.1\n",
|
||||
" center[(abs(center) < eps) & (center < 0)] = -eps\n",
|
||||
" center[(abs(center) < eps) & (center >= 0)] = eps\n",
|
||||
"\n",
|
||||
" return center\n",
|
||||
"def dist_per_sample(output, center):\n",
|
||||
" return tf.reduce_sum(tf.square(output - center), axis=-1)\n",
|
||||
"\n",
|
||||
"def score_per_sample(output, center, radius):\n",
|
||||
" return dist_per_sample(output, center) - radius**2\n",
|
||||
"\n",
|
||||
"def train_loss(output, center):\n",
|
||||
" return tf.reduce_mean(dist_per_sample(output, center))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "fd6f47c0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"center = get_center(deep_svdd_net, [X_face, X_eye])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b47b52f6",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def get_radius_from_arrays(nu, X_face, X_eye):\n",
|
||||
" z = deep_svdd_net.predict([X_face, X_eye])\n",
|
||||
" dists = dist_per_sample(z, center)\n",
|
||||
" return np.quantile(np.sqrt(dists), 1 - nu).astype(np.float32)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b062bd19",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"@tf.function\n",
|
||||
"def train_step(batch):\n",
|
||||
" with tf.GradientTape() as grad_tape:\n",
|
||||
" output = deep_svdd_net(batch, training=True)\n",
|
||||
" batch_loss = train_loss(output, center)\n",
|
||||
"\n",
|
||||
" gradients = grad_tape.gradient(batch_loss, deep_svdd_net.trainable_variables)\n",
|
||||
" optimizer.apply_gradients(zip(gradients, deep_svdd_net.trainable_variables))\n",
|
||||
"\n",
|
||||
" return batch_loss"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4c144130",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def train(dataset, epochs, nu):\n",
|
||||
" for epoch in range(epochs):\n",
|
||||
" start = time.time()\n",
|
||||
" losses = []\n",
|
||||
" for batch in dataset:\n",
|
||||
" batch_loss = train_step(batch)\n",
|
||||
" losses.append(batch_loss)\n",
|
||||
"\n",
|
||||
" print(f'{epoch+1}/{epochs} epoch: Loss of {np.mean(losses)} ({time.time()-start} secs)')\n",
|
||||
"\n",
|
||||
" return get_radius_from_arrays(nu, X_face, X_eye)\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"nu = 0.05 # Set nu respectively\n",
|
||||
"\n",
|
||||
"train_dataset = tf.data.Dataset.from_tensor_slices((X_face, X_eye)).shuffle(64).batch(64)\n",
|
||||
"\n",
|
||||
"optimizer = tf.keras.optimizers.Adam(1e-3)\n",
|
||||
"train(train_dataset, epochs=150, nu=nu)\n",
|
||||
"\n",
|
||||
"optimizer.learning_rate = 1e-4\n",
|
||||
"radius = train(train_dataset, 100, nu=nu)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "24f0cef0",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"prepare valid & test set"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "acb9c8f1",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Test set\n",
|
||||
"X_face_test = test_df_norm[face_au_cols].to_numpy(dtype=np.float32)\n",
|
||||
"X_eye_test = test_df_norm[eye_cols].to_numpy(dtype=np.float32)\n",
|
||||
"y_test = test_df_norm[\"label\"].to_numpy(dtype=np.float32)\n",
|
||||
"\n",
|
||||
"# Validation set\n",
|
||||
"X_face_val = val_df_norm[face_au_cols].to_numpy(dtype=np.float32)\n",
|
||||
"X_eye_val = val_df_norm[eye_cols].to_numpy(dtype=np.float32)\n",
|
||||
"y_val = val_df_norm[\"label\"].to_numpy(dtype=np.float32)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "49737d5d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"valid_scores = (score_per_sample(deep_svdd_net.predict([X_face_val, X_eye_val]), center, radius)).numpy()\n",
|
||||
"\n",
|
||||
"valid_fpr, valid_tpr, _ = roc_curve(y_val, valid_scores, pos_label=1)\n",
|
||||
"valid_auc = auc(valid_fpr, valid_tpr)\n",
|
||||
"\n",
|
||||
"plt.figure()\n",
|
||||
"plt.title('Deep SVDD')\n",
|
||||
"plt.plot(valid_fpr, valid_tpr, 'b-')\n",
|
||||
"plt.text(0.5, 0.5, f'AUC: {valid_auc:.4f}')\n",
|
||||
"plt.xlabel('False positive rate')\n",
|
||||
"plt.ylabel('True positive rate')\n",
|
||||
"plt.show()\n",
|
||||
"\n",
|
||||
"valid_predictions = (valid_scores > 0).astype(int)\n",
|
||||
"\n",
|
||||
"normal_acc = np.mean(valid_predictions[y_val == 0] == 0)\n",
|
||||
"anomaly_acc = np.mean(valid_predictions[y_val == 1] == 1)\n",
|
||||
"print(f'Accuracy on Validation set: {accuracy_score(y_val, valid_predictions)}')\n",
|
||||
"print(f'Accuracy for normals: {normal_acc:.4f}')\n",
|
||||
"print(f'Accuracy for anomalies: {anomaly_acc:.4f}')\n",
|
||||
"print(f'F1 on Validation set: {f1_score(y_val, valid_predictions)}')"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "475381db",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"deep_svdd_net.save(deep_svdd_save_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6ede1b15",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Results"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "c8481d07",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Validation set"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "719b41b2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"valid_predictions = (valid_scores > 0).astype(int)\n",
|
||||
"evaluation_tools.plot_confusion_matrix(true_labels=y_val, predictions=valid_predictions, label_names=[\"low\",\"high\"])\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "f33230b1",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Test set"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "f1189a28",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"test_scores = (\n",
|
||||
" score_per_sample(\n",
|
||||
" deep_svdd_net.predict([X_face_test, X_eye_test]),\n",
|
||||
" center,\n",
|
||||
" radius\n",
|
||||
" )\n",
|
||||
").numpy()\n",
|
||||
"\n",
|
||||
"test_predictions = (test_scores > 0).astype(int)\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "575dddcf",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"normal_acc = np.mean(test_predictions[y_test == 0] == 0)\n",
|
||||
"anomaly_acc = np.mean(test_predictions[y_test == 1] == 1)\n",
|
||||
"print(f'Accuracy on Test set: {accuracy_score(y_test, test_predictions)}')"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "5acade06",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"evaluation_tools.plot_confusion_matrix(true_labels=y_test, predictions=test_predictions, label_names=[\"low\",\"high\"])\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
@@ -111,18 +895,6 @@
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.12.10"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
|
||||
@@ -28,7 +28,7 @@
|
||||
"sys.path.append(base_dir)\n",
|
||||
"print(base_dir)\n",
|
||||
"\n",
|
||||
"from tools import evaluation_tools\n",
|
||||
"from Fahrsimulator_MSY2526_AI.model_training.tools import evaluation_tools, scaler\n",
|
||||
"from sklearn.preprocessing import StandardScaler, MinMaxScaler\n",
|
||||
"from sklearn.ensemble import IsolationForest\n",
|
||||
"from sklearn.model_selection import GridSearchCV, KFold\n",
|
||||
@@ -52,7 +52,7 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"data_path = Path(r\"C:\\Users\\micha\\FAUbox\\WS2526_Fahrsimulator_MSY (Celina Korzer)\\AU_dataset\\output_windowed.parquet\")"
|
||||
"data_path = Path(r\".parquet\") # TODO: set manually"
|
||||
]
|
||||
},
|
||||
{
|
||||
@@ -115,118 +115,6 @@
|
||||
"print(f\"high all: {high_all.shape}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "47a0f44d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def fit_normalizer(train_data, au_columns, method='standard', scope='global'):\n",
|
||||
" \"\"\"\n",
|
||||
" Fit normalization scalers on training data.\n",
|
||||
" \n",
|
||||
" Parameters:\n",
|
||||
" -----------\n",
|
||||
" train_data : pd.DataFrame\n",
|
||||
" Training dataframe with AU columns and subjectID\n",
|
||||
" au_columns : list\n",
|
||||
" List of AU column names to normalize\n",
|
||||
" method : str, default='standard'\n",
|
||||
" Normalization method: 'standard' for StandardScaler or 'minmax' for MinMaxScaler\n",
|
||||
" scope : str, default='global'\n",
|
||||
" Normalization scope: 'subject' for per-subject or 'global' for across all subjects\n",
|
||||
" \n",
|
||||
" Returns:\n",
|
||||
" --------\n",
|
||||
" dict\n",
|
||||
" Dictionary containing fitted scalers\n",
|
||||
" \"\"\"\n",
|
||||
" # Select scaler based on method\n",
|
||||
" if method == 'standard':\n",
|
||||
" Scaler = StandardScaler\n",
|
||||
" elif method == 'minmax':\n",
|
||||
" Scaler = MinMaxScaler\n",
|
||||
" else:\n",
|
||||
" raise ValueError(\"method must be 'standard' or 'minmax'\")\n",
|
||||
" \n",
|
||||
" scalers = {}\n",
|
||||
" \n",
|
||||
" if scope == 'subject':\n",
|
||||
" # Fit one scaler per subject\n",
|
||||
" for subject in train_data['subjectID'].unique():\n",
|
||||
" subject_mask = train_data['subjectID'] == subject\n",
|
||||
" scaler = Scaler()\n",
|
||||
" scaler.fit(train_data.loc[subject_mask, au_columns])\n",
|
||||
" scalers[subject] = scaler\n",
|
||||
" \n",
|
||||
" elif scope == 'global':\n",
|
||||
" # Fit one scaler for all subjects\n",
|
||||
" scaler = Scaler()\n",
|
||||
" scaler.fit(train_data[au_columns])\n",
|
||||
" scalers['global'] = scaler\n",
|
||||
" \n",
|
||||
" else:\n",
|
||||
" raise ValueError(\"scope must be 'subject' or 'global'\")\n",
|
||||
" \n",
|
||||
" return {'scalers': scalers, 'method': method, 'scope': scope}"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "642d0017",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def apply_normalizer(data, au_columns, normalizer_dict):\n",
|
||||
" \"\"\"\n",
|
||||
" Apply fitted normalization scalers to data.\n",
|
||||
" \n",
|
||||
" Parameters:\n",
|
||||
" -----------\n",
|
||||
" data : pd.DataFrame\n",
|
||||
" Dataframe with AU columns and subjectID\n",
|
||||
" au_columns : list\n",
|
||||
" List of AU column names to normalize\n",
|
||||
" normalizer_dict : dict\n",
|
||||
" Dictionary containing fitted scalers from fit_normalizer()\n",
|
||||
" \n",
|
||||
" Returns:\n",
|
||||
" --------\n",
|
||||
" pd.DataFrame\n",
|
||||
" DataFrame with normalized AU columns\n",
|
||||
" \"\"\"\n",
|
||||
" normalized_data = data.copy()\n",
|
||||
" scalers = normalizer_dict['scalers']\n",
|
||||
" scope = normalizer_dict['scope']\n",
|
||||
" \n",
|
||||
" if scope == 'subject':\n",
|
||||
" # Apply per-subject normalization\n",
|
||||
" for subject in data['subjectID'].unique():\n",
|
||||
" subject_mask = data['subjectID'] == subject\n",
|
||||
" \n",
|
||||
" # Use the subject's scaler if available, otherwise use a fitted scaler from training\n",
|
||||
" if subject in scalers:\n",
|
||||
" scaler = scalers[subject]\n",
|
||||
" else:\n",
|
||||
" # For new subjects not seen in training, use the first available scaler\n",
|
||||
" # (This is a fallback - ideally all test subjects should be in training for subject-level normalization)\n",
|
||||
" print(f\"Warning: Subject {subject} not found in training data. Using fallback scaler.\")\n",
|
||||
" scaler = list(scalers.values())[0]\n",
|
||||
" \n",
|
||||
" normalized_data.loc[subject_mask, au_columns] = scaler.transform(\n",
|
||||
" data.loc[subject_mask, au_columns]\n",
|
||||
" )\n",
|
||||
" \n",
|
||||
" elif scope == 'global':\n",
|
||||
" # Apply global normalization\n",
|
||||
" scaler = scalers['global']\n",
|
||||
" normalized_data[au_columns] = scaler.transform(data[au_columns])\n",
|
||||
" \n",
|
||||
" return normalized_data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "697b3cf7",
|
||||
@@ -301,20 +189,26 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Cell 2: Get AU columns and prepare datasets\n",
|
||||
"# Get all column names that start with 'AU'\n",
|
||||
"au_columns = [col for col in low_all.columns if col.startswith('AU')]\n",
|
||||
"au_columns = [col for col in low_all.columns if \"face\" in col.lower()] \n",
|
||||
"\n",
|
||||
"eye_columns = [ \n",
|
||||
" 'Fix_count_short_66_150','Fix_count_medium_300_500','Fix_count_long_gt_1000', \n",
|
||||
" 'Fix_count_100','Fix_mean_duration','Fix_median_duration', \n",
|
||||
" 'Sac_count','Sac_mean_amp','Sac_mean_dur','Sac_median_dur', \n",
|
||||
" 'Blink_count','Blink_mean_dur','Blink_median_dur', \n",
|
||||
" 'Pupil_mean','Pupil_IPA' \n",
|
||||
"] \n",
|
||||
"cols = au_columns +eye_columns\n",
|
||||
"# Prepare training data (only normal/low data)\n",
|
||||
"train_data = low_all[low_all['subjectID'].isin(train_subjects)][['subjectID'] + au_columns].copy()\n",
|
||||
"train_data = low_all[low_all['subjectID'].isin(train_subjects)][['subjectID'] + cols].copy()\n",
|
||||
"\n",
|
||||
"# Prepare validation data (normal and anomaly)\n",
|
||||
"val_normal_data = low_all[low_all['subjectID'].isin(val_subjects)][['subjectID'] + au_columns].copy()\n",
|
||||
"val_high_data = high_all[high_all['subjectID'].isin(val_subjects)][['subjectID'] + au_columns].copy()\n",
|
||||
"val_normal_data = low_all[low_all['subjectID'].isin(val_subjects)][['subjectID'] + cols].copy()\n",
|
||||
"val_high_data = high_all[high_all['subjectID'].isin(val_subjects)][['subjectID'] + cols].copy()\n",
|
||||
"\n",
|
||||
"# Prepare test data (normal and anomaly)\n",
|
||||
"test_normal_data = low_all[low_all['subjectID'].isin(test_subjects)][['subjectID'] + au_columns].copy()\n",
|
||||
"test_high_data = high_all[high_all['subjectID'].isin(test_subjects)][['subjectID'] + au_columns].copy()\n",
|
||||
"test_normal_data = low_all[low_all['subjectID'].isin(test_subjects)][['subjectID'] + cols].copy()\n",
|
||||
"test_high_data = high_all[high_all['subjectID'].isin(test_subjects)][['subjectID'] + cols].copy()\n",
|
||||
"\n",
|
||||
"print(f\"Train samples: {len(train_data)}\")\n",
|
||||
"print(f\"Val normal samples: {len(val_normal_data)}, Val high samples: {len(val_high_data)}\")\n",
|
||||
@@ -328,8 +222,8 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Cell 3: Fit normalizer on training data\n",
|
||||
"normalizer = fit_normalizer(train_data, au_columns, method='minmax', scope='global')\n",
|
||||
"# Fit normalizer on training data\n",
|
||||
"normalizer = scaler.fit_normalizer(train_data, cols, method='minmax', scope='global')\n",
|
||||
"print(\"Normalizer fitted on training data\")"
|
||||
]
|
||||
},
|
||||
@@ -340,12 +234,12 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Cell 4: Apply normalization to all datasets\n",
|
||||
"train_normalized = apply_normalizer(train_data, au_columns, normalizer)\n",
|
||||
"val_normal_normalized = apply_normalizer(val_normal_data, au_columns, normalizer)\n",
|
||||
"val_high_normalized = apply_normalizer(val_high_data, au_columns, normalizer)\n",
|
||||
"test_normal_normalized = apply_normalizer(test_normal_data, au_columns, normalizer)\n",
|
||||
"test_high_normalized = apply_normalizer(test_high_data, au_columns, normalizer)\n",
|
||||
"# Apply normalization to all datasets\n",
|
||||
"train_normalized = scaler.apply_normalizer(train_data, cols, normalizer)\n",
|
||||
"val_normal_normalized = scaler.apply_normalizer(val_normal_data, cols, normalizer)\n",
|
||||
"val_high_normalized = scaler.apply_normalizer(val_high_data, cols, normalizer)\n",
|
||||
"test_normal_normalized = scaler.apply_normalizer(test_normal_data, cols, normalizer)\n",
|
||||
"test_high_normalized = scaler.apply_normalizer(test_high_data, cols, normalizer)\n",
|
||||
"\n",
|
||||
"print(\"Normalization applied to all datasets\")"
|
||||
]
|
||||
@@ -357,11 +251,9 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Cell 5: Extract AU columns and create labels for grid search\n",
|
||||
"# Extract only AU columns (drop subjectID)\n",
|
||||
"X_train = train_normalized[au_columns].copy()\n",
|
||||
"X_val_normal = val_normal_normalized[au_columns].copy()\n",
|
||||
"X_val_high = val_high_normalized[au_columns].copy()\n",
|
||||
"X_train = train_normalized[cols].copy()\n",
|
||||
"X_val_normal = val_normal_normalized[cols].copy()\n",
|
||||
"X_val_high = val_high_normalized[cols].copy()\n",
|
||||
"\n",
|
||||
"# Combine train and validation sets for grid search\n",
|
||||
"X_grid_search = pd.concat([X_train, X_val_normal, X_val_high], ignore_index=True)\n",
|
||||
@@ -416,7 +308,7 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Cell 7: Train final model with best parameters on training data\n",
|
||||
"# Train final model with best parameters on training data\n",
|
||||
"final_model = IsolationForest(**best_params, random_state=42)\n",
|
||||
"final_model.fit(X_train.values)\n",
|
||||
"\n",
|
||||
@@ -430,9 +322,9 @@
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Cell 8: Prepare independent test set\n",
|
||||
"X_test_normal = test_normal_normalized[au_columns].copy()\n",
|
||||
"X_test_high = test_high_normalized[au_columns].copy()\n",
|
||||
"# Prepare independent test set\n",
|
||||
"X_test_normal = test_normal_normalized[cols].copy()\n",
|
||||
"X_test_high = test_high_normalized[cols].copy()\n",
|
||||
"\n",
|
||||
"# Combine test sets\n",
|
||||
"X_test = pd.concat([X_test_normal, X_test_high], ignore_index=True)\n",
|
||||
@@ -483,21 +375,9 @@
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "base",
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.11.5"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "e790b157",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Im folgenden wird auf die Daten das MAD Outlier removal angewendet."
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "46bd036d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import pandas as pd\n",
|
||||
"\n",
|
||||
"def calculate_mad_params(df, columns):\n",
|
||||
" \"\"\"\n",
|
||||
" Calculate median and MAD parameters for each column.\n",
|
||||
" This should be run ONLY on the training data.\n",
|
||||
" \n",
|
||||
" Returns a dictionary: {col: (median, mad)}\n",
|
||||
" \"\"\"\n",
|
||||
" params = {}\n",
|
||||
" for col in columns:\n",
|
||||
" median = df[col].median()\n",
|
||||
" mad = np.median(np.abs(df[col] - median))\n",
|
||||
" params[col] = (median, mad)\n",
|
||||
" return params"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e0691732",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def apply_mad_filter(df, params, threshold=3.5):\n",
|
||||
" \"\"\"\n",
|
||||
" Apply MAD-based outlier removal using precomputed parameters.\n",
|
||||
" Works on training, validation, and test data.\n",
|
||||
" \n",
|
||||
" df: DataFrame to filter\n",
|
||||
" params: dictionary {col: (median, mad)} from training data\n",
|
||||
" threshold: cutoff for robust Z-score\n",
|
||||
" \"\"\"\n",
|
||||
" df_clean = df.copy()\n",
|
||||
"\n",
|
||||
" for col, (median, mad) in params.items():\n",
|
||||
" if mad == 0:\n",
|
||||
" continue # no spread; nothing to remove for this column\n",
|
||||
"\n",
|
||||
" robust_z = 0.6745 * (df_clean[col] - median) / mad\n",
|
||||
" outlier_mask = np.abs(robust_z) > threshold\n",
|
||||
"\n",
|
||||
" # Remove values only in this specific column\n",
|
||||
" df_clean.loc[outlier_mask, col] = median\n",
|
||||
" print(df_clean.shape)\n",
|
||||
" \n",
|
||||
" print(df_clean.shape)\n",
|
||||
" return df_clean"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"language_info": {
|
||||
"name": "python"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,918 +0,0 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "708c9745",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Imports"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "53b10294",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"from pathlib import Path\n",
|
||||
"import sys\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"base_dir = os.path.abspath(os.path.join(os.getcwd(), \"..\"))\n",
|
||||
"sys.path.append(base_dir)\n",
|
||||
"print(base_dir)\n",
|
||||
"\n",
|
||||
"from Fahrsimulator_MSY2526_AI.model_training.tools import evaluation_tools, scaler, mad_outlier_removal\n",
|
||||
"from sklearn.preprocessing import StandardScaler, MinMaxScaler\n",
|
||||
"from sklearn.svm import OneClassSVM\n",
|
||||
"from sklearn.model_selection import GridSearchCV, KFold, ParameterGrid, train_test_split, GroupKFold\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import tensorflow as tf\n",
|
||||
"import pickle\n",
|
||||
"from sklearn.metrics import (roc_auc_score, accuracy_score, precision_score, recall_score, f1_score, confusion_matrix, classification_report, balanced_accuracy_score, ConfusionMatrixDisplay) "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "68101229",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### load Dataset"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "24a765e8",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dataset_path = Path(r\"/home/jovyan/data-paulusjafahrsimulator-gpu/first_AU_dataset/output_windowed.parquet\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "471001b0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.read_parquet(path=dataset_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "0fdecdaa",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Load Performance data and Subject Split"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "692d1b47",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"performance_path = Path(r\"/home/jovyan/data-paulusjafahrsimulator-gpu/subject_performance/3new_au_performance.csv\")\n",
|
||||
"performance_df = pd.read_csv(performance_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "ea617e3f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Subject IDs aus dem Haupt-Dataset nehmen\n",
|
||||
"subjects_from_df = df[\"subjectID\"].unique()\n",
|
||||
"\n",
|
||||
"# Performance-Subset nur für vorhandene Subjects\n",
|
||||
"perf_filtered = performance_df[\n",
|
||||
" performance_df[\"subjectID\"].isin(subjects_from_df)\n",
|
||||
"][[\"subjectID\", \"overall_score\"]]\n",
|
||||
"\n",
|
||||
"# Merge: nur Subjects, die sowohl im df als auch im Performance-CSV vorkommen\n",
|
||||
"merged = (\n",
|
||||
" pd.DataFrame({\"subjectID\": subjects_from_df})\n",
|
||||
" .merge(perf_filtered, on=\"subjectID\", how=\"inner\")\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Sicherstellen, dass keine Scores fehlen\n",
|
||||
"if merged[\"overall_score\"].isna().any():\n",
|
||||
" raise ValueError(\"Es fehlen Score-Werte für manche Subjects.\")\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "ae43df8d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"merged_sorted = merged.sort_values(\"overall_score\", ascending=False).reset_index(drop=True)\n",
|
||||
"\n",
|
||||
"scores = merged_sorted[\"overall_score\"].values\n",
|
||||
"n_total = len(merged_sorted)\n",
|
||||
"n_small = n_total // 3\n",
|
||||
"n_large = n_total - n_small\n",
|
||||
"\n",
|
||||
"# Schritt 1: zufällige Start-Aufteilung\n",
|
||||
"idx = np.arange(n_total)\n",
|
||||
"np.random.shuffle(idx)\n",
|
||||
"\n",
|
||||
"small_idx = idx[:n_small]\n",
|
||||
"large_idx = idx[n_small:]\n",
|
||||
"\n",
|
||||
"def score_diff(small_idx, large_idx):\n",
|
||||
" return abs(scores[small_idx].mean() - scores[large_idx].mean())\n",
|
||||
"\n",
|
||||
"diff = score_diff(small_idx, large_idx)\n",
|
||||
"threshold = 0.01\n",
|
||||
"max_iter = 100\n",
|
||||
"count = 0\n",
|
||||
"\n",
|
||||
"# Schritt 2: random swaps bis Differenz klein genug\n",
|
||||
"while diff > threshold and count < max_iter:\n",
|
||||
" # Zwei zufällige Elemente auswählen\n",
|
||||
" si = np.random.choice(small_idx)\n",
|
||||
" li = np.random.choice(large_idx)\n",
|
||||
" \n",
|
||||
" # Tausch durchführen\n",
|
||||
" new_small_idx = small_idx.copy()\n",
|
||||
" new_large_idx = large_idx.copy()\n",
|
||||
" \n",
|
||||
" new_small_idx[new_small_idx == si] = li\n",
|
||||
" new_large_idx[new_large_idx == li] = si\n",
|
||||
"\n",
|
||||
" # neue Differenz berechnen\n",
|
||||
" new_diff = score_diff(new_small_idx, new_large_idx)\n",
|
||||
"\n",
|
||||
" # Swap akzeptieren, wenn es besser wird\n",
|
||||
" if new_diff < diff:\n",
|
||||
" small_idx = new_small_idx\n",
|
||||
" large_idx = new_large_idx\n",
|
||||
" diff = new_diff\n",
|
||||
"\n",
|
||||
" count += 1\n",
|
||||
"\n",
|
||||
"# Finalgruppen\n",
|
||||
"group_small = merged_sorted.loc[small_idx].reset_index(drop=True)\n",
|
||||
"group_large = merged_sorted.loc[large_idx].reset_index(drop=True)\n",
|
||||
"\n",
|
||||
"print(\"Finale Score-Differenz:\", diff)\n",
|
||||
"print(\"Größe Gruppe 1:\", len(group_small))\n",
|
||||
"print(\"Größe Gruppe 2:\", len(group_large))\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9d1b414e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"group_large['overall_score'].mean()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "fa71f9a5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"group_small['overall_score'].mean()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "79ecb4a2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"training_subjects = group_large['subjectID'].values\n",
|
||||
"test_subjects = group_small['subjectID'].values\n",
|
||||
"print(training_subjects)\n",
|
||||
"print(test_subjects)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "87f9fe7d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"au_columns = [col for col in df.columns if col.lower().startswith(\"au\")]"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "009d268b",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Labeling"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4fa79163",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"low_all = df[\n",
|
||||
" ((df[\"PHASE\"] == \"baseline\") |\n",
|
||||
" ((df[\"STUDY\"] == \"n-back\") & (df[\"PHASE\"] != \"baseline\") & (df[\"LEVEL\"].isin([1, 4]))))\n",
|
||||
"]\n",
|
||||
"print(f\"low all: {low_all.shape}\")\n",
|
||||
"\n",
|
||||
"high_nback = df[\n",
|
||||
" (df[\"STUDY\"]==\"n-back\") &\n",
|
||||
" (df[\"LEVEL\"].isin([2, 3, 5, 6])) &\n",
|
||||
" (df[\"PHASE\"].isin([\"train\", \"test\"]))\n",
|
||||
"]\n",
|
||||
"print(f\"high n-back: {high_nback.shape}\")\n",
|
||||
"\n",
|
||||
"high_kdrive = df[\n",
|
||||
" (df[\"STUDY\"] == \"k-drive\") & (df[\"PHASE\"] != \"baseline\")\n",
|
||||
"]\n",
|
||||
"print(f\"high k-drive: {high_kdrive.shape}\")\n",
|
||||
"\n",
|
||||
"high_all = pd.concat([high_nback, high_kdrive])\n",
|
||||
"print(f\"high all: {high_all.shape}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "82b17d0b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"low = low_all.copy()\n",
|
||||
"high = high_all.copy()\n",
|
||||
"\n",
|
||||
"low[\"label\"] = 0\n",
|
||||
"high[\"label\"] = 1\n",
|
||||
"\n",
|
||||
"data = pd.concat([low, high], ignore_index=True)\n",
|
||||
"df = data.drop_duplicates()\n",
|
||||
"\n",
|
||||
"print(\"Label distribution:\")\n",
|
||||
"print(df[\"label\"].value_counts())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "4353f87c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Data cleaning with mad"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "c9afaf61",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# methode CT\n",
|
||||
"def calculate_mad_params(df, columns):\n",
|
||||
" \"\"\"\n",
|
||||
" Calculate median and MAD parameters for each column.\n",
|
||||
" This should be run ONLY on the training data.\n",
|
||||
" \n",
|
||||
" Returns a dictionary: {col: (median, mad)}\n",
|
||||
" \"\"\"\n",
|
||||
" params = {}\n",
|
||||
" for col in columns:\n",
|
||||
" median = df[col].median()\n",
|
||||
" mad = np.median(np.abs(df[col] - median))\n",
|
||||
" params[col] = (median, mad)\n",
|
||||
" return params\n",
|
||||
"\n",
|
||||
"def apply_mad_filter(df, params, threshold=3.5):\n",
|
||||
" \"\"\"\n",
|
||||
" Apply MAD-based outlier removal using precomputed parameters.\n",
|
||||
" Works on training, validation, and test data.\n",
|
||||
" \n",
|
||||
" df: DataFrame to filter\n",
|
||||
" params: dictionary {col: (median, mad)} from training data\n",
|
||||
" threshold: cutoff for robust Z-score\n",
|
||||
" \"\"\"\n",
|
||||
" df_clean = df.copy()\n",
|
||||
"\n",
|
||||
" for col, (median, mad) in params.items():\n",
|
||||
" if mad == 0:\n",
|
||||
" continue # no spread; nothing to remove for this column\n",
|
||||
"\n",
|
||||
" robust_z = 0.6745 * (df_clean[col] - median) / mad\n",
|
||||
" outlier_mask = np.abs(robust_z) > threshold\n",
|
||||
"\n",
|
||||
" # Remove values only in this specific column\n",
|
||||
" df_clean.loc[outlier_mask, col] = np.nan\n",
|
||||
" \n",
|
||||
" return df_clean"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4a286665",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_df = df[df.subjectID.isin(training_subjects)]\n",
|
||||
"test_df = df[df.subjectID.isin(test_subjects)]\n",
|
||||
"print(train_df.shape, test_df.shape)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "2671e0f4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"params = calculate_mad_params(train_df, au_columns)\n",
|
||||
"\n",
|
||||
"# Step 2: Apply filter consistently\n",
|
||||
"train_outlier_removed = apply_mad_filter(train_df, params, threshold=3.5)\n",
|
||||
"test_outlier_removed = apply_mad_filter(test_df, params, threshold=3.5)\n",
|
||||
"print(train_outlier_removed.shape, test_outlier_removed.shape)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "6c39b37f",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Normalisierung der Daten"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "5e6c654f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"normalizer = scaler.fit_normalizer(train_df, au_columns=au_columns, method='standard', scope='global')\n",
|
||||
"train_df_normal = scaler.apply_normalizer(train_df, au_columns=au_columns, normalizer_dict=normalizer)\n",
|
||||
"test_df_normal = scaler.apply_normalizer(test_df, au_columns=au_columns, normalizer_dict=normalizer)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "b6d25e7b",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"to do insert group k fold for train_df_normal"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "e826a998",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### AE first"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e6421371",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Beide Klassen für AE und SVM Training\n",
|
||||
"X_train_full = train_outlier_removed[au_columns].dropna()\n",
|
||||
"y_train_full = train_outlier_removed.loc[X_train_full.index, 'label'].values\n",
|
||||
"groups_train = train_outlier_removed.loc[X_train_full.index, 'subjectID'].values\n",
|
||||
"\n",
|
||||
"print(f\"Training data shape: {X_train_full.shape}\")\n",
|
||||
"print(f\"Label distribution in training: {pd.Series(y_train_full).value_counts()}\")\n",
|
||||
"\n",
|
||||
"# Test data\n",
|
||||
"X_test = test_outlier_removed[au_columns].dropna()\n",
|
||||
"y_test = test_outlier_removed.loc[X_test.index, 'label'].values\n",
|
||||
"\n",
|
||||
"print(f\"Test data shape: {X_test.shape}\")\n",
|
||||
"print(f\"Label distribution in test: {pd.Series(y_test).value_counts()}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "d982e47a",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Custom SVM Layer (differentiable approximation)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "50fbda1a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"class DifferentiableSVM(tf.keras.layers.Layer):\n",
|
||||
" \"\"\"\n",
|
||||
" Differentiable SVM Layer using hinge loss.\n",
|
||||
" This allows backpropagation through the SVM to the encoder.\n",
|
||||
" \"\"\"\n",
|
||||
" def __init__(self, C=1.0, **kwargs):\n",
|
||||
" super(DifferentiableSVM, self).__init__(**kwargs)\n",
|
||||
" self.C = C\n",
|
||||
" \n",
|
||||
" def build(self, input_shape):\n",
|
||||
" # SVM weights: w and bias b\n",
|
||||
" self.w = self.add_weight(\n",
|
||||
" shape=(input_shape[-1],),\n",
|
||||
" initializer='glorot_uniform',\n",
|
||||
" trainable=True,\n",
|
||||
" name='svm_w'\n",
|
||||
" )\n",
|
||||
" self.b = self.add_weight(\n",
|
||||
" shape=(1,),\n",
|
||||
" initializer='zeros',\n",
|
||||
" trainable=True,\n",
|
||||
" name='svm_b'\n",
|
||||
" )\n",
|
||||
" \n",
|
||||
" def call(self, inputs):\n",
|
||||
" # Decision function: w^T * x + b\n",
|
||||
" decision = tf.reduce_sum(inputs * self.w, axis=1, keepdims=True) + self.b\n",
|
||||
" return decision\n",
|
||||
" \n",
|
||||
" def compute_loss(self, inputs, labels):\n",
|
||||
" \"\"\"\n",
|
||||
" Hinge loss for SVM: max(0, 1 - y * (w^T * x + b))\n",
|
||||
" labels should be -1 or +1\n",
|
||||
" \"\"\"\n",
|
||||
" decision = self.call(inputs)\n",
|
||||
" \n",
|
||||
" # Convert labels from 0/1 to -1/+1\n",
|
||||
" labels_svm = tf.where(labels == 0, -1.0, 1.0)\n",
|
||||
" labels_svm = tf.cast(labels_svm, tf.float32)\n",
|
||||
" labels_svm = tf.reshape(labels_svm, (-1, 1))\n",
|
||||
" \n",
|
||||
" # Hinge loss\n",
|
||||
" hinge_loss = tf.reduce_mean(\n",
|
||||
" tf.maximum(0.0, 1.0 - labels_svm * decision)\n",
|
||||
" )\n",
|
||||
" \n",
|
||||
" # L2 regularization\n",
|
||||
" l2_loss = 0.5 * tf.reduce_sum(tf.square(self.w))\n",
|
||||
" \n",
|
||||
" return self.C * hinge_loss + l2_loss"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e7def811",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"class JointAESVM(tf.keras.Model):\n",
|
||||
" \"\"\"\n",
|
||||
" Joint Autoencoder + SVM Model\n",
|
||||
" Loss = reconstruction_loss + svm_loss\n",
|
||||
" \"\"\"\n",
|
||||
" def __init__(self, input_dim, latent_dim=5, hidden_dim=16, ae_weight=1.0, \n",
|
||||
" svm_weight=1.0, svm_C=1.0, reg=0.0001, **kwargs):\n",
|
||||
" super(JointAESVM, self).__init__(**kwargs)\n",
|
||||
" \n",
|
||||
" self.ae_weight = ae_weight\n",
|
||||
" self.svm_weight = svm_weight\n",
|
||||
" \n",
|
||||
" # Encoder\n",
|
||||
" self.encoder = tf.keras.Sequential([\n",
|
||||
" tf.keras.layers.Dense(input_dim, activation='relu', \n",
|
||||
" kernel_regularizer=tf.keras.regularizers.l2(reg)),\n",
|
||||
" tf.keras.layers.Dense(hidden_dim, activation='relu',\n",
|
||||
" kernel_regularizer=tf.keras.regularizers.l2(reg)),\n",
|
||||
" tf.keras.layers.Dense(latent_dim, activation='relu',\n",
|
||||
" kernel_regularizer=tf.keras.regularizers.l2(reg))\n",
|
||||
" ], name='encoder')\n",
|
||||
" \n",
|
||||
" # Decoder\n",
|
||||
" self.decoder = tf.keras.Sequential([\n",
|
||||
" tf.keras.layers.Dense(latent_dim, activation='relu',\n",
|
||||
" kernel_regularizer=tf.keras.regularizers.l2(reg)),\n",
|
||||
" tf.keras.layers.Dense(hidden_dim, activation='relu',\n",
|
||||
" kernel_regularizer=tf.keras.regularizers.l2(reg)),\n",
|
||||
" tf.keras.layers.Dense(input_dim, activation='linear',\n",
|
||||
" kernel_regularizer=tf.keras.regularizers.l2(reg))\n",
|
||||
" ], name='decoder')\n",
|
||||
" \n",
|
||||
" # SVM Layer\n",
|
||||
" self.svm = DifferentiableSVM(C=svm_C, name='svm')\n",
|
||||
" \n",
|
||||
" def call(self, inputs, training=False):\n",
|
||||
" # Encode\n",
|
||||
" encoded = self.encoder(inputs, training=training)\n",
|
||||
" \n",
|
||||
" # Decode (for reconstruction)\n",
|
||||
" decoded = self.decoder(encoded, training=training)\n",
|
||||
" \n",
|
||||
" # SVM decision (for classification)\n",
|
||||
" svm_output = self.svm(encoded)\n",
|
||||
" \n",
|
||||
" return decoded, svm_output, encoded\n",
|
||||
" \n",
|
||||
" def compute_loss(self, x, y_true):\n",
|
||||
" # Forward pass\n",
|
||||
" x_reconstructed, svm_decision, encoded = self(x, training=True)\n",
|
||||
" \n",
|
||||
" # Reconstruction loss (MSE)\n",
|
||||
" reconstruction_loss = tf.reduce_mean(\n",
|
||||
" tf.square(x - x_reconstructed)\n",
|
||||
" )\n",
|
||||
" \n",
|
||||
" # SVM loss (hinge)\n",
|
||||
" svm_loss = self.svm.compute_loss(encoded, y_true)\n",
|
||||
" \n",
|
||||
" # Total loss\n",
|
||||
" total_loss = (self.ae_weight * reconstruction_loss + \n",
|
||||
" self.svm_weight * svm_loss)\n",
|
||||
" \n",
|
||||
" return total_loss, reconstruction_loss, svm_loss\n",
|
||||
"\n",
|
||||
"print(\"Joint AE-SVM Model class defined\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "541085f3",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Train function"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "d0bf18e3",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def train_joint_model(X_train, y_train, groups, model_params, \n",
|
||||
" epochs=200, batch_size=64, learning_rate=0.0001):\n",
|
||||
" \"\"\"\n",
|
||||
" Train joint model on given data\n",
|
||||
" \"\"\"\n",
|
||||
" # Build model\n",
|
||||
" model = JointAESVM(\n",
|
||||
" input_dim=X_train.shape[1],\n",
|
||||
" latent_dim=model_params['latent_dim'],\n",
|
||||
" hidden_dim=model_params['hidden_dim'],\n",
|
||||
" ae_weight=model_params['ae_weight'],\n",
|
||||
" svm_weight=model_params['svm_weight'],\n",
|
||||
" svm_C=model_params['svm_C'],\n",
|
||||
" reg=model_params['reg']\n",
|
||||
" )\n",
|
||||
" \n",
|
||||
" optimizer = tf.keras.optimizers.Adam(learning_rate=learning_rate)\n",
|
||||
" \n",
|
||||
" # Training history\n",
|
||||
" history = {\n",
|
||||
" 'total_loss': [],\n",
|
||||
" 'recon_loss': [],\n",
|
||||
" 'svm_loss': []\n",
|
||||
" }\n",
|
||||
" \n",
|
||||
" # Convert to tensors\n",
|
||||
" X_train_tf = tf.constant(X_train.values, dtype=tf.float32)\n",
|
||||
" y_train_tf = tf.constant(y_train, dtype=tf.float32)\n",
|
||||
" \n",
|
||||
" # Create dataset\n",
|
||||
" dataset = tf.data.Dataset.from_tensor_slices((X_train_tf, y_train_tf))\n",
|
||||
" dataset = dataset.shuffle(buffer_size=1024).batch(batch_size)\n",
|
||||
" \n",
|
||||
" # Training loop\n",
|
||||
" for epoch in range(epochs):\n",
|
||||
" epoch_loss = 0.0\n",
|
||||
" epoch_recon = 0.0\n",
|
||||
" epoch_svm = 0.0\n",
|
||||
" n_batches = 0\n",
|
||||
" \n",
|
||||
" for x_batch, y_batch in dataset:\n",
|
||||
" with tf.GradientTape() as tape:\n",
|
||||
" total_loss, recon_loss, svm_loss = model.compute_loss(x_batch, y_batch)\n",
|
||||
" \n",
|
||||
" # Backpropagation\n",
|
||||
" gradients = tape.gradient(total_loss, model.trainable_variables)\n",
|
||||
" optimizer.apply_gradients(zip(gradients, model.trainable_variables))\n",
|
||||
" \n",
|
||||
" epoch_loss += total_loss.numpy()\n",
|
||||
" epoch_recon += recon_loss.numpy()\n",
|
||||
" epoch_svm += svm_loss.numpy()\n",
|
||||
" n_batches += 1\n",
|
||||
" \n",
|
||||
" # Average losses\n",
|
||||
" history['total_loss'].append(epoch_loss / n_batches)\n",
|
||||
" history['recon_loss'].append(epoch_recon / n_batches)\n",
|
||||
" history['svm_loss'].append(epoch_svm / n_batches)\n",
|
||||
" \n",
|
||||
" if (epoch + 1) % 20 == 0:\n",
|
||||
" print(f\"Epoch {epoch+1}/{epochs} - \"\n",
|
||||
" f\"Total: {history['total_loss'][-1]:.4f}, \"\n",
|
||||
" f\"Recon: {history['recon_loss'][-1]:.4f}, \"\n",
|
||||
" f\"SVM: {history['svm_loss'][-1]:.4f}\")\n",
|
||||
" \n",
|
||||
" return model, history\n",
|
||||
"\n",
|
||||
"print(\"Training function defined\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b6a04540",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Parameter Grid\n",
|
||||
"param_grid = {\n",
|
||||
" 'latent_dim': [5, 8],\n",
|
||||
" 'hidden_dim': [10, 16],\n",
|
||||
" 'ae_weight': [0.5, 1.0],\n",
|
||||
" 'svm_weight': [0.5, 1.0, 2.0],\n",
|
||||
" 'svm_C': [0.1, 1.0, 10.0],\n",
|
||||
" 'reg': [0.0001, 0.001]\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"n_splits = 5 # Weniger Splits wegen Rechenzeit\n",
|
||||
"gkf = GroupKFold(n_splits=n_splits)\n",
|
||||
"\n",
|
||||
"print(f\"Starting Grid Search with {n_splits}-fold GroupKFold\")\n",
|
||||
"print(f\"Parameter combinations: {len(list(ParameterGrid(param_grid)))}\")\n",
|
||||
"print(\"This will take a while...\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "228463ce",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def evaluate_model(model, X, y):\n",
|
||||
" \"\"\"Evaluate joint model\"\"\"\n",
|
||||
" X_tf = tf.constant(X, dtype=tf.float32)\n",
|
||||
" _, svm_decision, _ = model(X_tf, training=False)\n",
|
||||
" \n",
|
||||
" # Predict: decision > 0 -> class 1, else class 0\n",
|
||||
" y_pred = (svm_decision.numpy().flatten() > 0).astype(int)\n",
|
||||
" \n",
|
||||
" bal_accuracy = balanced_accuracy_score(y, y_pred)\n",
|
||||
" return bal_accuracy, y_pred\n",
|
||||
"\n",
|
||||
"print(\"Evaluation function defined\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "c945fc87",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Grid Search\n",
|
||||
"best_score = -np.inf\n",
|
||||
"best_params = None\n",
|
||||
"best_model = None\n",
|
||||
"all_results = []\n",
|
||||
"\n",
|
||||
"X_train_array = X_train_full.values\n",
|
||||
"y_train_array = y_train_full\n",
|
||||
"\n",
|
||||
"for param_idx, params in enumerate(ParameterGrid(param_grid)):\n",
|
||||
" print(f\"\\n{'='*60}\")\n",
|
||||
" print(f\"Testing parameters {param_idx + 1}/{len(list(ParameterGrid(param_grid)))}\")\n",
|
||||
" print(f\"Params: {params}\")\n",
|
||||
" print(f\"{'='*60}\")\n",
|
||||
" \n",
|
||||
" fold_scores = []\n",
|
||||
" \n",
|
||||
" for fold, (train_idx, val_idx) in enumerate(gkf.split(X_train_array, y_train_array, groups_train)):\n",
|
||||
" print(f\"\\nFold {fold + 1}/{n_splits}\")\n",
|
||||
" \n",
|
||||
" X_fold_train = pd.DataFrame(X_train_array[train_idx], columns=X_train_full.columns)\n",
|
||||
" y_fold_train = y_train_array[train_idx]\n",
|
||||
" X_fold_val = X_train_array[val_idx]\n",
|
||||
" y_fold_val = y_train_array[val_idx]\n",
|
||||
" \n",
|
||||
" # Train model\n",
|
||||
" model, history = train_joint_model(\n",
|
||||
" X_fold_train, y_fold_train, groups_train[train_idx],\n",
|
||||
" model_params=params,\n",
|
||||
" epochs=100, # Weniger Epochen für Grid Search\n",
|
||||
" batch_size=64,\n",
|
||||
" learning_rate=0.0001\n",
|
||||
" )\n",
|
||||
" \n",
|
||||
" # Validate\n",
|
||||
" val_bal_acc, _ = evaluate_model(model, X_fold_val, y_fold_val)\n",
|
||||
" fold_scores.append(val_bal_acc)\n",
|
||||
" print(f\"Fold {fold + 1} Validation balanced Accuracy: {val_bal_acc:.4f}\")\n",
|
||||
" \n",
|
||||
" mean_score = np.mean(fold_scores)\n",
|
||||
" std_score = np.std(fold_scores)\n",
|
||||
" \n",
|
||||
" result = {\n",
|
||||
" **params,\n",
|
||||
" 'mean_cv_bal_accuracy': mean_score,\n",
|
||||
" 'std_cv_bal_accuracy': std_score\n",
|
||||
" }\n",
|
||||
" all_results.append(result)\n",
|
||||
" \n",
|
||||
" print(f\"\\nMean CV bal. Accuracy: {mean_score:.4f} ± {std_score:.4f}\")\n",
|
||||
" \n",
|
||||
" if mean_score > best_score:\n",
|
||||
" best_score = mean_score\n",
|
||||
" best_params = params\n",
|
||||
" print(\"*** NEW BEST PARAMETERS ***\")\n",
|
||||
"\n",
|
||||
"print(f\"\\n{'='*60}\")\n",
|
||||
"print(\"GRID SEARCH COMPLETED\")\n",
|
||||
"print(f\"{'='*60}\")\n",
|
||||
"print(f\"Best parameters: {best_params}\")\n",
|
||||
"print(f\"Best CV bal. accuracy: {best_score:.4f}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "0a0606f5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"results_df = pd.DataFrame(all_results)\n",
|
||||
"results_df = results_df.sort_values('mean_cv_accuracy', ascending=False)\n",
|
||||
"\n",
|
||||
"print(\"\\nTop 10 configurations:\")\n",
|
||||
"print(results_df.head(10))\n",
|
||||
"\n",
|
||||
"# Plot\n",
|
||||
"plt.figure(figsize=(12, 6))\n",
|
||||
"plt.barh(range(min(10, len(results_df))), \n",
|
||||
" results_df['mean_cv_accuracy'].head(10))\n",
|
||||
"plt.yticks(range(min(10, len(results_df))), \n",
|
||||
" [f\"Config {i+1}\" for i in range(min(10, len(results_df)))])\n",
|
||||
"plt.xlabel('Mean CV Accuracy')\n",
|
||||
"plt.title('Top 10 Configurations')\n",
|
||||
"plt.tight_layout()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "87906b05",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"print(\"Training final model on all training data...\")\n",
|
||||
"print(f\"Best parameters: {best_params}\")\n",
|
||||
"\n",
|
||||
"final_model, final_history = train_joint_model(\n",
|
||||
" X_train_full, y_train_full, groups_train,\n",
|
||||
" model_params=best_params,\n",
|
||||
" epochs=300, # Mehr Epochen für finales Training\n",
|
||||
" batch_size=64,\n",
|
||||
" learning_rate=0.0001\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"print(\"\\nFinal model training completed!\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "718137a8",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"fig, axes = plt.subplots(1, 3, figsize=(15, 4))\n",
|
||||
"\n",
|
||||
"axes[0].plot(final_history['total_loss'])\n",
|
||||
"axes[0].set_title('Total Loss')\n",
|
||||
"axes[0].set_xlabel('Epoch')\n",
|
||||
"axes[0].set_ylabel('Loss')\n",
|
||||
"axes[0].grid(True, alpha=0.3)\n",
|
||||
"\n",
|
||||
"axes[1].plot(final_history['recon_loss'])\n",
|
||||
"axes[1].set_title('Reconstruction Loss')\n",
|
||||
"axes[1].set_xlabel('Epoch')\n",
|
||||
"axes[1].set_ylabel('Loss')\n",
|
||||
"axes[1].grid(True, alpha=0.3)\n",
|
||||
"\n",
|
||||
"axes[2].plot(final_history['svm_loss'])\n",
|
||||
"axes[2].set_title('SVM Loss')\n",
|
||||
"axes[2].set_xlabel('Epoch')\n",
|
||||
"axes[2].set_ylabel('Loss')\n",
|
||||
"axes[2].grid(True, alpha=0.3)\n",
|
||||
"\n",
|
||||
"plt.tight_layout()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "02fbc5a2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Get predictions\n",
|
||||
"test_acc, y_pred = evaluate_model(final_model, X_test.values, y_test)\n",
|
||||
"\n",
|
||||
"# Get SVM decision values for ROC-AUC\n",
|
||||
"X_test_tf = tf.constant(X_test.values, dtype=tf.float32)\n",
|
||||
"_, svm_decision, _ = final_model(X_test_tf, training=False)\n",
|
||||
"y_pred_decision = svm_decision.numpy().flatten()\n",
|
||||
"\n",
|
||||
"# Metrics\n",
|
||||
"print(\"=\" * 50)\n",
|
||||
"print(\"TEST SET EVALUATION\")\n",
|
||||
"print(\"=\" * 50)\n",
|
||||
"print(f\"\\nAccuracy: {accuracy_score(y_test, y_pred):.4f}\")\n",
|
||||
"print(f\"Precision: {precision_score(y_test, y_pred):.4f}\")\n",
|
||||
"print(f\"Recall: {recall_score(y_test, y_pred):.4f}\")\n",
|
||||
"print(f\"F1-Score: {f1_score(y_test, y_pred):.4f}\")\n",
|
||||
"\n",
|
||||
"# ROC-AUC (decision values as probability proxy)\n",
|
||||
"decision_scaled = MinMaxScaler().fit_transform(y_pred_decision.reshape(-1, 1)).flatten()\n",
|
||||
"print(f\"ROC-AUC: {roc_auc_score(y_test, decision_scaled):.4f}\")\n",
|
||||
"\n",
|
||||
"print(\"\\nConfusion Matrix:\")\n",
|
||||
"cm = confusion_matrix(y_test, y_pred)\n",
|
||||
"print(cm)\n",
|
||||
"\n",
|
||||
"print(\"\\nClassification Report:\")\n",
|
||||
"print(classification_report(y_test, y_pred))\n",
|
||||
"\n",
|
||||
"# Visualize Confusion Matrix\n",
|
||||
"fig, ax = plt.subplots(figsize=(8, 6))\n",
|
||||
"disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=['Low Load (0)', 'High Load (1)'])\n",
|
||||
"disp.plot(cmap='Blues', ax=ax, colorbar=True, values_format='d')\n",
|
||||
"ax.set_title('Confusion Matrix - Test Set', fontsize=14, fontweight='bold')\n",
|
||||
"plt.tight_layout()\n",
|
||||
"plt.show()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4c524bce",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Save entire model\n",
|
||||
"final_model.save_weights('joint_ae_svm_weights.h5')\n",
|
||||
"print(\"Model weights saved as 'joint_ae_svm_weights.h5'\")\n",
|
||||
"\n",
|
||||
"# Save encoder separately\n",
|
||||
"final_model.encoder.save('encoder_joint.keras')\n",
|
||||
"print(\"Encoder saved as 'encoder_joint.keras'\")\n",
|
||||
"\n",
|
||||
"# Save best parameters\n",
|
||||
"with open('best_params_joint.pkl', 'wb') as f:\n",
|
||||
" pickle.dump(best_params, f)\n",
|
||||
"print(\"Best parameters saved as 'best_params_joint.pkl'\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "792c658d",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"* doch mal svm ae pipeline?\n",
|
||||
"* einfach mal mit 20 13 5\n",
|
||||
"* label hinzufügen\n",
|
||||
"* mad von CT verwenden oder wert anpassen, ggf. vergleich welches label wie oft vorkommt vorher und nachher. --> labelling schritt von CT übernehmen\n"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -1,254 +0,0 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "708c9745",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Imports"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "53b10294",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"from pathlib import Path\n",
|
||||
"import sys\n",
|
||||
"import os\n",
|
||||
"\n",
|
||||
"base_dir = os.path.abspath(os.path.join(os.getcwd(), \"..\"))\n",
|
||||
"sys.path.append(base_dir)\n",
|
||||
"print(base_dir)\n",
|
||||
"\n",
|
||||
"from Fahrsimulator_MSY2526_AI.model_training.tools import evaluation_tools, scaler, mad_outlier_removal\n",
|
||||
"from sklearn.preprocessing import StandardScaler, MinMaxScaler\n",
|
||||
"from sklearn.svm import OneClassSVM\n",
|
||||
"from sklearn.model_selection import GridSearchCV, KFold, ParameterGrid, train_test_split\n",
|
||||
"import matplotlib.pyplot as plt\n",
|
||||
"import tensorflow as tf\n",
|
||||
"import pickle\n",
|
||||
"from sklearn.metrics import (roc_auc_score, accuracy_score, precision_score, \n",
|
||||
" recall_score, f1_score, confusion_matrix, classification_report) "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "68101229",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### load Dataset"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "24a765e8",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dataset_path = Path(r\"/home/jovyan/data-paulusjafahrsimulator-gpu/first_AU_dataset/output_windowed.parquet\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "471001b0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.read_parquet(path=dataset_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "0fdecdaa",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Load Performance data and Subject Split"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "692d1b47",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"performance_path = Path(r\"/home/jovyan/data-paulusjafahrsimulator-gpu/subject_performance/3new_au_performance.csv\")\n",
|
||||
"performance_df = pd.read_csv(performance_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "ea617e3f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Subject IDs aus dem Haupt-Dataset nehmen\n",
|
||||
"subjects_from_df = df[\"subjectID\"].unique()\n",
|
||||
"\n",
|
||||
"# Performance-Subset nur für vorhandene Subjects\n",
|
||||
"perf_filtered = performance_df[\n",
|
||||
" performance_df[\"subjectID\"].isin(subjects_from_df)\n",
|
||||
"][[\"subjectID\", \"overall_score\"]]\n",
|
||||
"\n",
|
||||
"# Merge: nur Subjects, die sowohl im df als auch im Performance-CSV vorkommen\n",
|
||||
"merged = (\n",
|
||||
" pd.DataFrame({\"subjectID\": subjects_from_df})\n",
|
||||
" .merge(perf_filtered, on=\"subjectID\", how=\"inner\")\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Sicherstellen, dass keine Scores fehlen\n",
|
||||
"if merged[\"overall_score\"].isna().any():\n",
|
||||
" raise ValueError(\"Es fehlen Score-Werte für manche Subjects.\")\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "ae43df8d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"merged_sorted = merged.sort_values(\"overall_score\", ascending=False).reset_index(drop=True)\n",
|
||||
"\n",
|
||||
"scores = merged_sorted[\"overall_score\"].values\n",
|
||||
"n_total = len(merged_sorted)\n",
|
||||
"n_small = n_total // 3\n",
|
||||
"n_large = n_total - n_small\n",
|
||||
"\n",
|
||||
"# Schritt 1: zufällige Start-Aufteilung\n",
|
||||
"idx = np.arange(n_total)\n",
|
||||
"np.random.shuffle(idx)\n",
|
||||
"\n",
|
||||
"small_idx = idx[:n_small]\n",
|
||||
"large_idx = idx[n_small:]\n",
|
||||
"\n",
|
||||
"def score_diff(small_idx, large_idx):\n",
|
||||
" return abs(scores[small_idx].mean() - scores[large_idx].mean())\n",
|
||||
"\n",
|
||||
"diff = score_diff(small_idx, large_idx)\n",
|
||||
"threshold = 0.01\n",
|
||||
"max_iter = 100\n",
|
||||
"count = 0\n",
|
||||
"\n",
|
||||
"# Schritt 2: random swaps bis Differenz klein genug\n",
|
||||
"while diff > threshold and count < max_iter:\n",
|
||||
" # Zwei zufällige Elemente auswählen\n",
|
||||
" si = np.random.choice(small_idx)\n",
|
||||
" li = np.random.choice(large_idx)\n",
|
||||
" \n",
|
||||
" # Tausch durchführen\n",
|
||||
" new_small_idx = small_idx.copy()\n",
|
||||
" new_large_idx = large_idx.copy()\n",
|
||||
" \n",
|
||||
" new_small_idx[new_small_idx == si] = li\n",
|
||||
" new_large_idx[new_large_idx == li] = si\n",
|
||||
"\n",
|
||||
" # neue Differenz berechnen\n",
|
||||
" new_diff = score_diff(new_small_idx, new_large_idx)\n",
|
||||
"\n",
|
||||
" # Swap akzeptieren, wenn es besser wird\n",
|
||||
" if new_diff < diff:\n",
|
||||
" small_idx = new_small_idx\n",
|
||||
" large_idx = new_large_idx\n",
|
||||
" diff = new_diff\n",
|
||||
"\n",
|
||||
" count += 1\n",
|
||||
"\n",
|
||||
"# Finalgruppen\n",
|
||||
"group_small = merged_sorted.loc[small_idx].reset_index(drop=True)\n",
|
||||
"group_large = merged_sorted.loc[large_idx].reset_index(drop=True)\n",
|
||||
"\n",
|
||||
"print(\"Finale Score-Differenz:\", diff)\n",
|
||||
"print(\"Größe Gruppe 1:\", len(group_small))\n",
|
||||
"print(\"Größe Gruppe 2:\", len(group_large))\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "9d1b414e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"group_large['overall_score'].mean()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "fa71f9a5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"group_small['overall_score'].mean()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "79ecb4a2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"training_subjects = group_large['subjectID'].values\n",
|
||||
"test_subjects = group_small['subjectID'].values\n",
|
||||
"print(training_subjects)\n",
|
||||
"print(test_subjects)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "4353f87c",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"### Data cleaning with mad"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "76610052",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# SET\n",
|
||||
"threshold_mad = 100\n",
|
||||
"column_praefix ='AU'\n",
|
||||
"\n",
|
||||
"au_columns = [col for col in df.columns if col.startswith(column_praefix)]\n",
|
||||
"cleaned_df = mad_outlier_removal(df,columns=au_columns, threshold=threshold_mad)\n",
|
||||
"print(cleaned_df.shape)\n",
|
||||
"print(df.shape)"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.12.10"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -21,3 +21,42 @@ def mad_outlier_removal(df, columns, threshold=3.5, c=1.4826):
|
||||
|
||||
final_mask = np.logical_and.reduce(masks)
|
||||
return df_clean[final_mask]
|
||||
|
||||
def calculate_mad_params(df, columns):
|
||||
"""
|
||||
Calculate median and MAD parameters for each column.
|
||||
This should be run ONLY on the training data.
|
||||
|
||||
Returns a dictionary: {col: (median, mad)}
|
||||
"""
|
||||
params = {}
|
||||
for col in columns:
|
||||
median = df[col].median()
|
||||
mad = np.median(np.abs(df[col] - median))
|
||||
params[col] = (median, mad)
|
||||
return params
|
||||
|
||||
def apply_mad_filter(df, params, threshold=3.5):
|
||||
"""
|
||||
Apply MAD-based outlier removal using precomputed parameters.
|
||||
Works on training, validation, and test data.
|
||||
|
||||
df: DataFrame to filter
|
||||
params: dictionary {col: (median, mad)} from training data
|
||||
threshold: cutoff for robust Z-score
|
||||
"""
|
||||
df_clean = df.copy()
|
||||
|
||||
for col, (median, mad) in params.items():
|
||||
if mad == 0:
|
||||
continue # no spread; nothing to remove for this column
|
||||
|
||||
robust_z = 0.6745 * (df_clean[col] - median) / mad
|
||||
outlier_mask = np.abs(robust_z) > threshold
|
||||
|
||||
# Remove values only in this specific column
|
||||
df_clean.loc[outlier_mask, col] = median
|
||||
|
||||
|
||||
print(df_clean.shape)
|
||||
return df_clean
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
from sklearn.preprocessing import MinMaxScaler, StandardScaler
|
||||
import pandas as pd
|
||||
import pickle
|
||||
from sklearn.preprocessing import StandardScaler, MinMaxScaler
|
||||
import numpy as np
|
||||
import os
|
||||
|
||||
def fit_normalizer(train_data, au_columns, method='standard', scope='global'):
|
||||
"""
|
||||
@@ -19,9 +21,8 @@ def fit_normalizer(train_data, au_columns, method='standard', scope='global'):
|
||||
Returns:
|
||||
--------
|
||||
dict
|
||||
Dictionary containing fitted scalers
|
||||
Dictionary containing fitted scalers and statistics for new subjects
|
||||
"""
|
||||
# Select scaler based on method
|
||||
if method == 'standard':
|
||||
Scaler = StandardScaler
|
||||
elif method == 'minmax':
|
||||
@@ -30,19 +31,54 @@ def fit_normalizer(train_data, au_columns, method='standard', scope='global'):
|
||||
raise ValueError("method must be 'standard' or 'minmax'")
|
||||
|
||||
scalers = {}
|
||||
|
||||
if scope == 'subject':
|
||||
# Fit one scaler per subject
|
||||
subject_stats = []
|
||||
|
||||
for subject in train_data['subjectID'].unique():
|
||||
subject_mask = train_data['subjectID'] == subject
|
||||
scaler = Scaler()
|
||||
scaler.fit(train_data.loc[subject_mask, au_columns])
|
||||
scaler.fit(train_data.loc[subject_mask, au_columns].values)
|
||||
scalers[subject] = scaler
|
||||
|
||||
# Store statistics for averaging
|
||||
if method == 'standard':
|
||||
subject_stats.append({
|
||||
'mean': scaler.mean_,
|
||||
'std': scaler.scale_
|
||||
})
|
||||
elif method == 'minmax':
|
||||
subject_stats.append({
|
||||
'min': scaler.data_min_,
|
||||
'max': scaler.data_max_
|
||||
})
|
||||
|
||||
# Calculate average statistics for new subjects
|
||||
if method == 'standard':
|
||||
avg_mean = np.mean([s['mean'] for s in subject_stats], axis=0)
|
||||
avg_std = np.mean([s['std'] for s in subject_stats], axis=0)
|
||||
fallback_scaler = StandardScaler()
|
||||
fallback_scaler.mean_ = avg_mean
|
||||
fallback_scaler.scale_ = avg_std
|
||||
fallback_scaler.var_ = avg_std ** 2
|
||||
fallback_scaler.n_features_in_ = len(au_columns)
|
||||
elif method == 'minmax':
|
||||
avg_min = np.mean([s['min'] for s in subject_stats], axis=0)
|
||||
avg_max = np.mean([s['max'] for s in subject_stats], axis=0)
|
||||
fallback_scaler = MinMaxScaler()
|
||||
fallback_scaler.data_min_ = avg_min
|
||||
fallback_scaler.data_max_ = avg_max
|
||||
fallback_scaler.data_range_ = avg_max - avg_min
|
||||
fallback_scaler.scale_ = 1.0 / fallback_scaler.data_range_
|
||||
fallback_scaler.min_ = -avg_min * fallback_scaler.scale_
|
||||
fallback_scaler.n_features_in_ = len(au_columns)
|
||||
|
||||
scalers['_fallback'] = fallback_scaler
|
||||
|
||||
elif scope == 'global':
|
||||
# Fit one scaler for all subjects
|
||||
scaler = Scaler()
|
||||
scaler.fit(train_data[au_columns])
|
||||
scaler.fit(train_data[au_columns].values)
|
||||
scalers['global'] = scaler
|
||||
|
||||
else:
|
||||
@@ -50,7 +86,7 @@ def fit_normalizer(train_data, au_columns, method='standard', scope='global'):
|
||||
|
||||
return {'scalers': scalers, 'method': method, 'scope': scope}
|
||||
|
||||
def apply_normalizer(data, au_columns, normalizer_dict):
|
||||
def apply_normalizer(data, columns, normalizer_dict):
|
||||
"""
|
||||
Apply fitted normalization scalers to data.
|
||||
|
||||
@@ -71,28 +107,70 @@ def apply_normalizer(data, au_columns, normalizer_dict):
|
||||
normalized_data = data.copy()
|
||||
scalers = normalizer_dict['scalers']
|
||||
scope = normalizer_dict['scope']
|
||||
normalized_data[columns] = normalized_data[columns].astype(np.float64)
|
||||
|
||||
if scope == 'subject':
|
||||
# Apply per-subject normalization
|
||||
for subject in data['subjectID'].unique():
|
||||
subject_mask = data['subjectID'] == subject
|
||||
|
||||
# Use the subject's scaler if available, otherwise use a fitted scaler from training
|
||||
# Use the subject's scaler if available, otherwise use fallback
|
||||
if subject in scalers:
|
||||
scaler = scalers[subject]
|
||||
else:
|
||||
# For new subjects not seen in training, use the first available scaler
|
||||
# (This is a fallback - ideally all test subjects should be in training for subject-level normalization)
|
||||
print(f"Warning: Subject {subject} not found in training data. Using fallback scaler.")
|
||||
scaler = list(scalers.values())[0]
|
||||
# Use averaged scaler for new subjects
|
||||
scaler = scalers['_fallback']
|
||||
print(f"Info: Subject {subject} not in training data. Using averaged scaler from training subjects.")
|
||||
|
||||
normalized_data.loc[subject_mask, au_columns] = scaler.transform(
|
||||
data.loc[subject_mask, au_columns]
|
||||
normalized_data.loc[subject_mask, columns] = scaler.transform(
|
||||
data.loc[subject_mask, columns].values
|
||||
)
|
||||
|
||||
elif scope == 'global':
|
||||
# Apply global normalization
|
||||
scaler = scalers['global']
|
||||
normalized_data[au_columns] = scaler.transform(data[au_columns])
|
||||
normalized_data[columns] = scaler.transform(data[columns].values)
|
||||
|
||||
return normalized_data
|
||||
|
||||
|
||||
|
||||
def save_normalizer(normalizer_dict, filepath):
|
||||
"""
|
||||
Save fitted normalizer to disk.
|
||||
|
||||
Parameters:
|
||||
-----------
|
||||
normalizer_dict : dict
|
||||
Dictionary containing fitted scalers from fit_normalizer()
|
||||
filepath : str
|
||||
Path to save the normalizer (e.g., 'normalizer.pkl')
|
||||
"""
|
||||
# Create directory if it does not exist
|
||||
dirpath = os.path.dirname(filepath)
|
||||
if dirpath:
|
||||
os.makedirs(dirpath, exist_ok=True)
|
||||
|
||||
with open(filepath, 'wb') as f:
|
||||
pickle.dump(normalizer_dict, f)
|
||||
|
||||
print(f"Normalizer saved to {filepath}")
|
||||
|
||||
def load_normalizer(filepath):
|
||||
"""
|
||||
Load fitted normalizer from disk.
|
||||
|
||||
Parameters:
|
||||
-----------
|
||||
filepath : str
|
||||
Path to the saved normalizer file
|
||||
|
||||
Returns:
|
||||
--------
|
||||
dict
|
||||
Dictionary containing fitted scalers
|
||||
"""
|
||||
with open(filepath, 'rb') as f:
|
||||
normalizer_dict = pickle.load(f)
|
||||
print(f"Normalizer loaded from {filepath}")
|
||||
return normalizer_dict
|
||||
@@ -0,0 +1,807 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e3be057e-8d2a-4d05-bd42-6b1dc75df5ed",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd\n",
|
||||
"from pathlib import Path\n",
|
||||
"from sklearn.preprocessing import StandardScaler, MinMaxScaler"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "13ad96f5",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# data_path = Path(r\"~/Fahrsimulator_MSY2526_AI/model_training/xgboost/output_windowed.parquet\")\n",
|
||||
"data_path = Path(r\"~/data-paulusjafahrsimulator-gpu/new_datasets/combined_dataset_25hz.parquet\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4aa1e32c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import pandas as pd\n",
|
||||
"import numpy as np\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"def performance_based_split(\n",
|
||||
" subject_ids,\n",
|
||||
" performance_df,\n",
|
||||
" split_ratio=0.33,\n",
|
||||
" threshold=0.01,\n",
|
||||
" max_iter=100,\n",
|
||||
" random_seed=None\n",
|
||||
"):\n",
|
||||
" \"\"\"\n",
|
||||
" Split subjects into two groups based on performance scores with balanced means.\n",
|
||||
" \n",
|
||||
" Parameters\n",
|
||||
" ----------\n",
|
||||
" subject_ids : array-like\n",
|
||||
" List or array of subject IDs present in your dataset\n",
|
||||
" performance_df : pd.DataFrame\n",
|
||||
" DataFrame containing 'subjectID' and 'overall_score' columns\n",
|
||||
" split_ratio : float, default=0.33\n",
|
||||
" Proportion of subjects for the smaller group (0 < split_ratio < 1)\n",
|
||||
" threshold : float, default=0.01\n",
|
||||
" Target difference threshold between group means\n",
|
||||
" max_iter : int, default=100\n",
|
||||
" Maximum number of swap iterations\n",
|
||||
" random_seed : int, optional\n",
|
||||
" Random seed for reproducibility\n",
|
||||
" \n",
|
||||
" Returns\n",
|
||||
" -------\n",
|
||||
" group_small_ids : np.ndarray\n",
|
||||
" Subject IDs for the smaller group\n",
|
||||
" group_large_ids : np.ndarray\n",
|
||||
" Subject IDs for the larger group\n",
|
||||
" score_diff : float\n",
|
||||
" Final absolute difference between group means\n",
|
||||
" \n",
|
||||
" Raises\n",
|
||||
" ------\n",
|
||||
" ValueError\n",
|
||||
" If subjects are missing performance scores or no subjects match\n",
|
||||
" \"\"\"\n",
|
||||
" if random_seed is not None:\n",
|
||||
" np.random.seed(random_seed)\n",
|
||||
" \n",
|
||||
" # Filter performance data\n",
|
||||
" perf_filtered = performance_df[\n",
|
||||
" performance_df[\"subjectID\"].isin(subject_ids)\n",
|
||||
" ][[\"subjectID\", \"overall_score\"]]\n",
|
||||
" \n",
|
||||
" # Merge to get only subjects present in both dataset and performance file\n",
|
||||
" merged = (\n",
|
||||
" pd.DataFrame({\"subjectID\": subject_ids})\n",
|
||||
" .merge(perf_filtered, on=\"subjectID\", how=\"inner\")\n",
|
||||
" )\n",
|
||||
" \n",
|
||||
" if len(merged) == 0:\n",
|
||||
" raise ValueError(\"No subjects found in both dataset and performance file.\")\n",
|
||||
" \n",
|
||||
" # Check for missing scores\n",
|
||||
" if merged[\"overall_score\"].isna().any():\n",
|
||||
" raise ValueError(\"Missing score values for some subjects.\")\n",
|
||||
" \n",
|
||||
" merged_sorted = merged.sort_values(\"overall_score\", ascending=False).reset_index(drop=True)\n",
|
||||
" \n",
|
||||
" scores = merged_sorted[\"overall_score\"].values\n",
|
||||
" n_total = len(merged_sorted)\n",
|
||||
" n_small = int(n_total * split_ratio)\n",
|
||||
" n_large = n_total - n_small\n",
|
||||
" \n",
|
||||
" # Initial random split\n",
|
||||
" idx = np.arange(n_total)\n",
|
||||
" np.random.shuffle(idx)\n",
|
||||
" \n",
|
||||
" small_idx = idx[:n_small]\n",
|
||||
" large_idx = idx[n_small:]\n",
|
||||
" \n",
|
||||
" def score_diff(small_idx, large_idx):\n",
|
||||
" return abs(scores[small_idx].mean() - scores[large_idx].mean())\n",
|
||||
" \n",
|
||||
" diff = score_diff(small_idx, large_idx)\n",
|
||||
" count = 0\n",
|
||||
" \n",
|
||||
" # Optimize via random swaps\n",
|
||||
" while diff > threshold and count < max_iter:\n",
|
||||
" si = np.random.choice(small_idx)\n",
|
||||
" li = np.random.choice(large_idx)\n",
|
||||
" \n",
|
||||
" new_small_idx = small_idx.copy()\n",
|
||||
" new_large_idx = large_idx.copy()\n",
|
||||
" \n",
|
||||
" new_small_idx[new_small_idx == si] = li\n",
|
||||
" new_large_idx[new_large_idx == li] = si\n",
|
||||
" \n",
|
||||
" new_diff = score_diff(new_small_idx, new_large_idx)\n",
|
||||
" \n",
|
||||
" if new_diff < diff:\n",
|
||||
" small_idx = new_small_idx\n",
|
||||
" large_idx = new_large_idx\n",
|
||||
" diff = new_diff\n",
|
||||
" \n",
|
||||
" count += 1\n",
|
||||
" \n",
|
||||
" # Extract subject IDs\n",
|
||||
" group_small_ids = merged_sorted.loc[small_idx, \"subjectID\"].values\n",
|
||||
" group_large_ids = merged_sorted.loc[large_idx, \"subjectID\"].values\n",
|
||||
" \n",
|
||||
" return group_small_ids, group_large_ids, diff"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "95e1a351",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df = pd.read_parquet(path=data_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "248d519b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"performance_path = Path(r\"/home/jovyan/data-paulusjafahrsimulator-gpu/subject_performance/3new_au_performance.csv\")\n",
|
||||
"performance_df = pd.read_csv(performance_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "8b9992e0",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_ids, temp_ids, diff1 = performance_based_split(\n",
|
||||
" subject_ids=df[\"subjectID\"].unique(),\n",
|
||||
" performance_df=performance_df,\n",
|
||||
" split_ratio=0.6, # 60% train, 40% temp\n",
|
||||
" random_seed=42\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"val_ids, test_ids, diff2 = performance_based_split(\n",
|
||||
" subject_ids=temp_ids,\n",
|
||||
" performance_df=performance_df,\n",
|
||||
" split_ratio=0.5, # 50/50 split of remaining 40%\n",
|
||||
" random_seed=43\n",
|
||||
")\n",
|
||||
"print(diff1, diff2)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "68afd83e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"subjects = df['subjectID'].unique()\n",
|
||||
"print(subjects)\n",
|
||||
"print(len(subjects))\n",
|
||||
"print(len(subjects)*0.66)\n",
|
||||
"print(len(subjects)*0.33)\n",
|
||||
"print(df.columns)\n",
|
||||
"print(df['STUDY'].unique())\n",
|
||||
"print(df['LEVEL'].unique())\n",
|
||||
"print(df['PHASE'].unique())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "52dfd885",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"low_all = df[\n",
|
||||
" ((df[\"PHASE\"] == \"baseline\") |\n",
|
||||
" ((df[\"STUDY\"] == \"n-back\") & (df[\"PHASE\"] != \"baseline\") & (df[\"LEVEL\"].isin([1, 4]))))\n",
|
||||
"]\n",
|
||||
"print(f\"low all: {low_all.shape}\")\n",
|
||||
"\n",
|
||||
"high_nback = df[\n",
|
||||
" (df[\"STUDY\"]==\"n-back\") &\n",
|
||||
" (df[\"LEVEL\"].isin([2, 3, 5, 6])) &\n",
|
||||
" (df[\"PHASE\"].isin([\"train\", \"test\"]))\n",
|
||||
"]\n",
|
||||
"print(f\"high n-back: {high_nback.shape}\")\n",
|
||||
"\n",
|
||||
"high_kdrive = df[\n",
|
||||
" (df[\"STUDY\"] == \"k-drive\") & (df[\"PHASE\"] != \"baseline\")\n",
|
||||
"]\n",
|
||||
"print(f\"high k-drive: {high_kdrive.shape}\")\n",
|
||||
"\n",
|
||||
"high_all = pd.concat([high_nback, high_kdrive])\n",
|
||||
"print(f\"high all: {high_all.shape}\")"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "8fba6edf",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from sklearn.preprocessing import MinMaxScaler, StandardScaler\n",
|
||||
"import pandas as pd\n",
|
||||
"\n",
|
||||
"def fit_normalizer(train_data, au_columns, method='standard', scope='global'):\n",
|
||||
" \"\"\"\n",
|
||||
" Fit normalization scalers on training data.\n",
|
||||
" \n",
|
||||
" Parameters:\n",
|
||||
" -----------\n",
|
||||
" train_data : pd.DataFrame\n",
|
||||
" Training dataframe with AU columns and subjectID\n",
|
||||
" au_columns : list\n",
|
||||
" List of AU column names to normalize\n",
|
||||
" method : str, default='standard'\n",
|
||||
" Normalization method: 'standard' for StandardScaler or 'minmax' for MinMaxScaler\n",
|
||||
" scope : str, default='global'\n",
|
||||
" Normalization scope: 'subject' for per-subject or 'global' for across all subjects\n",
|
||||
" \n",
|
||||
" Returns:\n",
|
||||
" --------\n",
|
||||
" dict\n",
|
||||
" Dictionary containing fitted scalers\n",
|
||||
" \"\"\"\n",
|
||||
" # Select scaler based on method\n",
|
||||
" if method == 'standard':\n",
|
||||
" Scaler = StandardScaler\n",
|
||||
" elif method == 'minmax':\n",
|
||||
" Scaler = MinMaxScaler\n",
|
||||
" else:\n",
|
||||
" raise ValueError(\"method must be 'standard' or 'minmax'\")\n",
|
||||
" \n",
|
||||
" scalers = {}\n",
|
||||
" \n",
|
||||
" if scope == 'subject':\n",
|
||||
" # Fit one scaler per subject\n",
|
||||
" for subject in train_data['subjectID'].unique():\n",
|
||||
" subject_mask = train_data['subjectID'] == subject\n",
|
||||
" scaler = Scaler()\n",
|
||||
" scaler.fit(train_data.loc[subject_mask, au_columns])\n",
|
||||
" scalers[subject] = scaler\n",
|
||||
" \n",
|
||||
" elif scope == 'global':\n",
|
||||
" # Fit one scaler for all subjects\n",
|
||||
" scaler = Scaler()\n",
|
||||
" scaler.fit(train_data[au_columns])\n",
|
||||
" scalers['global'] = scaler\n",
|
||||
" \n",
|
||||
" else:\n",
|
||||
" raise ValueError(\"scope must be 'subject' or 'global'\")\n",
|
||||
" \n",
|
||||
" return {'scalers': scalers, 'method': method, 'scope': scope}\n",
|
||||
"\n",
|
||||
"def apply_normalizer(data, au_columns, normalizer_dict):\n",
|
||||
" \"\"\"\n",
|
||||
" Apply fitted normalization scalers to data.\n",
|
||||
" \n",
|
||||
" Parameters:\n",
|
||||
" -----------\n",
|
||||
" data : pd.DataFrame\n",
|
||||
" Dataframe with AU columns and subjectID\n",
|
||||
" au_columns : list\n",
|
||||
" List of AU column names to normalize\n",
|
||||
" normalizer_dict : dict\n",
|
||||
" Dictionary containing fitted scalers from fit_normalizer()\n",
|
||||
" \n",
|
||||
" Returns:\n",
|
||||
" --------\n",
|
||||
" pd.DataFrame\n",
|
||||
" DataFrame with normalized AU columns\n",
|
||||
" \"\"\"\n",
|
||||
" normalized_data = data.copy()\n",
|
||||
" scalers = normalizer_dict['scalers']\n",
|
||||
" scope = normalizer_dict['scope']\n",
|
||||
" \n",
|
||||
" if scope == 'subject':\n",
|
||||
" # Apply per-subject normalization\n",
|
||||
" for subject in data['subjectID'].unique():\n",
|
||||
" subject_mask = data['subjectID'] == subject\n",
|
||||
" \n",
|
||||
" # Use the subject's scaler if available, otherwise use a fitted scaler from training\n",
|
||||
" if subject in scalers:\n",
|
||||
" scaler = scalers[subject]\n",
|
||||
" else:\n",
|
||||
" # For new subjects not seen in training, use the first available scaler\n",
|
||||
" # (This is a fallback - ideally all test subjects should be in training for subject-level normalization)\n",
|
||||
" print(f\"Warning: Subject {subject} not found in training data. Using fallback scaler.\")\n",
|
||||
" scaler = list(scalers.values())[0]\n",
|
||||
" \n",
|
||||
" normalized_data.loc[subject_mask, au_columns] = scaler.transform(\n",
|
||||
" data.loc[subject_mask, au_columns]\n",
|
||||
" )\n",
|
||||
" \n",
|
||||
" elif scope == 'global':\n",
|
||||
" # Apply global normalization\n",
|
||||
" scaler = scalers['global']\n",
|
||||
" normalized_data[au_columns] = scaler.transform(data[au_columns])\n",
|
||||
" \n",
|
||||
" return normalized_data"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "24e3a77b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"%pip install xgboost"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "8e7fa0fa",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"from sklearn.model_selection import train_test_split,StratifiedKFold, GridSearchCV\n",
|
||||
"from sklearn.metrics import accuracy_score, f1_score, roc_auc_score, classification_report, confusion_matrix\n",
|
||||
"import xgboost as xgb\n",
|
||||
"import joblib\n",
|
||||
"import matplotlib.pyplot as plt"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "325ef71c",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"low = low_all.copy()\n",
|
||||
"high = high_all.copy()\n",
|
||||
"\n",
|
||||
"low[\"label\"] = 0\n",
|
||||
"high[\"label\"] = 1\n",
|
||||
"\n",
|
||||
"data = pd.concat([low, high], ignore_index=True)\n",
|
||||
"data = data.drop_duplicates()\n",
|
||||
"\n",
|
||||
"print(\"Label distribution:\")\n",
|
||||
"print(data[\"label\"].value_counts())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "67d70e84",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"face_au_cols = [c for c in train_df.columns if c.startswith(\"FACE_AU\")]\n",
|
||||
"eye_cols = ['Fix_count_short_66_150', 'Fix_count_medium_300_500',\n",
|
||||
" 'Fix_count_long_gt_1000', 'Fix_count_100', 'Fix_mean_duration',\n",
|
||||
" 'Fix_median_duration', 'Sac_count', 'Sac_mean_amp', 'Sac_mean_dur',\n",
|
||||
" 'Sac_median_dur', 'Blink_count', 'Blink_mean_dur', 'Blink_median_dur',\n",
|
||||
" 'Pupil_mean', 'Pupil_IPA']\n",
|
||||
"print(len(eye_cols))\n",
|
||||
"all_signal_columns = face_au_cols+eye_cols\n",
|
||||
"print(len(all_signal_columns))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b19eb87b",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"low_all = df[\n",
|
||||
" ((df[\"PHASE\"] == \"baseline\") |\n",
|
||||
" ((df[\"STUDY\"] == \"n-back\") & (df[\"PHASE\"] != \"baseline\") & (df[\"LEVEL\"].isin([1, 4]))))\n",
|
||||
"]\n",
|
||||
"print(f\"low all: {low_all.shape}\")\n",
|
||||
"\n",
|
||||
"high_nback = df[\n",
|
||||
" (df[\"STUDY\"]==\"n-back\") &\n",
|
||||
" (df[\"LEVEL\"].isin([2, 3, 5, 6])) &\n",
|
||||
" (df[\"PHASE\"].isin([\"train\", \"test\"]))\n",
|
||||
"]\n",
|
||||
"print(f\"high n-back: {high_nback.shape}\")\n",
|
||||
"\n",
|
||||
"high_kdrive = df[\n",
|
||||
" (df[\"STUDY\"] == \"k-drive\") & (df[\"PHASE\"] != \"baseline\")\n",
|
||||
"]\n",
|
||||
"print(f\"high k-drive: {high_kdrive.shape}\")\n",
|
||||
"\n",
|
||||
"high_all = pd.concat([high_nback, high_kdrive])\n",
|
||||
"print(f\"high all: {high_all.shape}\")\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"\n",
|
||||
"low = low_all.copy()\n",
|
||||
"high = high_all.copy()\n",
|
||||
"\n",
|
||||
"low[\"label\"] = 0\n",
|
||||
"high[\"label\"] = 1\n",
|
||||
"\n",
|
||||
"data = pd.concat([low, high], ignore_index=True)\n",
|
||||
"df = data.drop_duplicates()\n",
|
||||
"\n",
|
||||
"print(\"Label distribution:\")\n",
|
||||
"print(df[\"label\"].value_counts())"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "960bb8c7",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"train_df = df[\n",
|
||||
" (df.subjectID.isin(train_ids)) & (df['label'] == 0)\n",
|
||||
"].copy()\n",
|
||||
"\n",
|
||||
"# Validation: balanced sampling of label=0 and label=1\n",
|
||||
"val_df_full = df[df.subjectID.isin(val_ids)].copy()\n",
|
||||
"\n",
|
||||
"# Get all label=0 samples\n",
|
||||
"val_df_label0 = val_df_full[val_df_full['label'] == 0]\n",
|
||||
"\n",
|
||||
"# Sample same number from label=1\n",
|
||||
"n_samples = len(val_df_label0)\n",
|
||||
"val_df_label1 = val_df_full[val_df_full['label'] == 1].sample(\n",
|
||||
" n=n_samples, random_state=42\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Combine\n",
|
||||
"val_df = pd.concat([val_df_label0, val_df_label1], ignore_index=True)\n",
|
||||
"test_df = df[df.subjectID.isin(test_ids)]\n",
|
||||
"print(train_df.shape, val_df.shape,test_df.shape)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "dbb58abd",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import numpy as np\n",
|
||||
"import pandas as pd\n",
|
||||
"\n",
|
||||
"def calculate_mad_params(df, columns):\n",
|
||||
" \"\"\"\n",
|
||||
" Calculate median and MAD parameters for each column.\n",
|
||||
" This should be run ONLY on the training data.\n",
|
||||
" \n",
|
||||
" Returns a dictionary: {col: (median, mad)}\n",
|
||||
" \"\"\"\n",
|
||||
" params = {}\n",
|
||||
" for col in columns:\n",
|
||||
" median = df[col].median()\n",
|
||||
" mad = np.median(np.abs(df[col] - median))\n",
|
||||
" params[col] = (median, mad)\n",
|
||||
" return params\n",
|
||||
"\n",
|
||||
"def apply_mad_filter(df, params, threshold=3.5):\n",
|
||||
" \"\"\"\n",
|
||||
" Apply MAD-based outlier removal using precomputed parameters.\n",
|
||||
" Works on training, validation, and test data.\n",
|
||||
" \n",
|
||||
" df: DataFrame to filter\n",
|
||||
" params: dictionary {col: (median, mad)} from training data\n",
|
||||
" threshold: cutoff for robust Z-score\n",
|
||||
" \"\"\"\n",
|
||||
" df_clean = df.copy()\n",
|
||||
"\n",
|
||||
" for col, (median, mad) in params.items():\n",
|
||||
" if mad == 0:\n",
|
||||
" continue # no spread; nothing to remove for this column\n",
|
||||
"\n",
|
||||
" robust_z = 0.6745 * (df_clean[col] - median) / mad\n",
|
||||
" outlier_mask = np.abs(robust_z) > threshold\n",
|
||||
"\n",
|
||||
" # Remove values only in this specific column\n",
|
||||
" df_clean.loc[outlier_mask, col] = median\n",
|
||||
" \n",
|
||||
" return df_clean"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "0f03f1b4",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# # Step 1: Fit parameters on training data\n",
|
||||
"# params = calculate_mad_params(train_df, au_columns)\n",
|
||||
"\n",
|
||||
"# # Step 2: Apply filter consistently\n",
|
||||
"# train_outlier_removed = apply_mad_filter(train_df, params, threshold=3.5)\n",
|
||||
"# val_outlier_removed = apply_mad_filter(val_df, params, threshold=50)\n",
|
||||
"# test_outlier_removed = apply_mad_filter(test_df, params, threshold=50)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "289f6b89",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"print(train_df.subjectID.unique())\n",
|
||||
"print(df.subjectID.unique())\n",
|
||||
"\n",
|
||||
"normalizer = fit_normalizer(df, all_signal_columns, method='standard', scope='subject')\n",
|
||||
"train_df_norm = apply_normalizer(train_df, all_signal_columns, normalizer)\n",
|
||||
"val_df_norm = apply_normalizer(val_df, all_signal_columns, normalizer)\n",
|
||||
"test_df_norm = apply_normalizer(test_df, all_signal_columns, normalizer)\n",
|
||||
"\n",
|
||||
"# normalizer = fit_normalizer(train_outlier_removed, au_columns, method=\"standard\", scope=\"global\")\n",
|
||||
"\n",
|
||||
"# train_scaled = apply_normalizer(train_outlier_removed, normalizer, au_columns)\n",
|
||||
"# val_scaled = apply_normalizer(val_df, normalizer, au_columns)\n",
|
||||
"# test_scaled = apply_normalizer(test_df, normalizer, au_columns)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "5df30e8d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"X_train, y_train = train_df[all_signal_columns].values, train_df[\"label\"].values\n",
|
||||
"X_val, y_val = val_df[all_signal_columns].values, val_df[\"label\"].values\n",
|
||||
"X_test, y_test = test_df[all_signal_columns].values, test_df[\"label\"].values"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "6fb7c86a",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import xgboost as xgb\n",
|
||||
"from sklearn.model_selection import GroupKFold, GridSearchCV\n",
|
||||
"\n",
|
||||
"# EarlyStopping mit kürzerem Patience\n",
|
||||
"early_stop = xgb.callback.EarlyStopping(\n",
|
||||
" rounds=25, metric_name='auc', data_name='validation_0', save_best=True\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Basis-Modell: nur feste Parameter, keine Optimierungswerte\n",
|
||||
"xgb_clf = xgb.XGBClassifier(\n",
|
||||
" objective=\"binary:logistic\",\n",
|
||||
" scale_pos_weight=1100/1550, # Klassenungleichgewicht berücksichtigen\n",
|
||||
" eval_metric=[\"logloss\", \"auc\", \"error\"],\n",
|
||||
" use_label_encoder=False,\n",
|
||||
" random_state=42,\n",
|
||||
" callbacks=[early_stop],\n",
|
||||
" verbosity=0\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Parameter-Raster für GridSearch\n",
|
||||
"param_grid = {\n",
|
||||
" \"learning_rate\": [0.01, 0.05, 0.1],\n",
|
||||
" \"max_depth\": [2, 3],\n",
|
||||
" \"subsample\": [0.5, 0.6, 0.7],\n",
|
||||
" \"colsample_bytree\": [0.5, 0.6, 0.7],\n",
|
||||
" \"reg_alpha\": [0.1, 1, 5, 10],\n",
|
||||
" \"reg_lambda\": [5, 10, 20, 50],\n",
|
||||
" \"min_child_weight\": [10, 20, 50],\n",
|
||||
" \"max_delta_step\": [1, 5, 10],\n",
|
||||
" \"n_estimators\": [500, 1000, 2000]\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"# K-Fold Cross Validation\n",
|
||||
"cv = GroupKFold(n_splits=5, shuffle=True, random_state=42)\n",
|
||||
"\n",
|
||||
"# Grid Search Setup\n",
|
||||
"grid_search = GridSearchCV(\n",
|
||||
" estimator=xgb_clf,\n",
|
||||
" param_grid=param_grid,\n",
|
||||
" scoring=\"roc_auc\",\n",
|
||||
" n_jobs=-1,\n",
|
||||
" cv=cv,\n",
|
||||
" verbose=2\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"# Training mit Cross Validation, Gruppen übergeben\n",
|
||||
"X_train = train_df[all_signal_columns].values\n",
|
||||
"y_train = train_df[\"label\"].values\n",
|
||||
"groups = train_df[\"subjectID\"].values\n",
|
||||
"\n",
|
||||
"# Training mit Cross Validation\n",
|
||||
"grid_search.fit(\n",
|
||||
" X_train, y_train,\n",
|
||||
" groups=groups,\n",
|
||||
" eval_set=[(X_train, y_train), (X_val, y_val)],\n",
|
||||
" verbose=False,\n",
|
||||
")\n",
|
||||
"\n",
|
||||
"print(\"Beste Parameter:\", grid_search.best_params_)\n",
|
||||
"print(\"Bestes AUC:\", grid_search.best_score_)\n",
|
||||
"\n",
|
||||
"# Bestes Modell extrahieren\n",
|
||||
"model = grid_search.best_estimator_"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "d2681022",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# Plots\n",
|
||||
"\n",
|
||||
"results = model.evals_result()\n",
|
||||
"epochs = len(results['validation_0']['auc'])\n",
|
||||
"x_axis = range(0, epochs)\n",
|
||||
"\n",
|
||||
"# --- Plot Loss ---\n",
|
||||
"plt.figure(figsize=(8,6))\n",
|
||||
"plt.plot(x_axis, results['validation_0']['logloss'], label='Validation Loss')\n",
|
||||
"plt.plot(x_axis, results['validation_1']['logloss'], label='Training Loss')\n",
|
||||
"plt.legend()\n",
|
||||
"plt.xlabel('Epochs')\n",
|
||||
"plt.ylabel('Logloss')\n",
|
||||
"plt.title('XGBoost Loss during Training')\n",
|
||||
"plt.grid(True)\n",
|
||||
"plt.show()\n",
|
||||
"\n",
|
||||
"# --- Plot Accuracy ---\n",
|
||||
"plt.figure(figsize=(8,6))\n",
|
||||
"plt.plot(x_axis, [1-e for e in results['validation_0']['error']], label='Validation Accuracy')\n",
|
||||
"plt.plot(x_axis, [1-e for e in results['validation_1']['error']], label='Training Accuracy')\n",
|
||||
"plt.legend()\n",
|
||||
"plt.xlabel('Epochs')\n",
|
||||
"plt.ylabel('Accuracy')\n",
|
||||
"plt.title('XGBoost Accuracy during Training')\n",
|
||||
"plt.grid(True)\n",
|
||||
"plt.show()\n",
|
||||
"\n",
|
||||
"# Plot AUC\n",
|
||||
"\n",
|
||||
"plt.figure(figsize=(8,6))\n",
|
||||
"plt.plot(x_axis, results['validation_0']['auc'], label='Validation AUC')\n",
|
||||
"plt.plot(x_axis, results['validation_1']['auc'], marker='o')\n",
|
||||
"plt.legend()\n",
|
||||
"plt.xlabel('Epochs')\n",
|
||||
"plt.ylabel('AUC')\n",
|
||||
"plt.title('XGBoost AUC during Training')\n",
|
||||
"plt.grid(True)\n",
|
||||
"plt.show()\n",
|
||||
"\n",
|
||||
"# ROC-Kurve plotten\n",
|
||||
"y_pred_proba = model.predict_proba(X_val)[:, 1]\n",
|
||||
"# RocCurveDisplay.from_predictions(y_val, y_pred_proba)\n",
|
||||
"plt.title(\"ROC Curve (Validation Set)\")\n",
|
||||
"plt.grid(True)\n",
|
||||
"plt.show()\n",
|
||||
"\n",
|
||||
"# Test: Loss und Accuracy\n",
|
||||
"y_test_proba = model.predict_proba(X_test)[:,1]\n",
|
||||
"y_test_pred = (y_test_proba > 0.5).astype(int)\n",
|
||||
"\n",
|
||||
"# print(\"Test Loss:\", log_loss(y_test, y_test_proba))\n",
|
||||
"print(\"Test Accuracy:\", accuracy_score(y_test, y_test_pred))"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "09a8cd21",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"from sklearn.metrics import confusion_matrix, accuracy_score, f1_score, roc_auc_score, classification_report, ConfusionMatrixDisplay\n",
|
||||
"\n",
|
||||
"def evaluate(model, X, y, title=\"Evaluation\"):\n",
|
||||
" # Vorhersagen\n",
|
||||
" preds_proba = model.predict_proba(X)[:, 1]\n",
|
||||
" preds = (preds_proba > 0.5).astype(int)\n",
|
||||
"\n",
|
||||
" # Metriken ausgeben\n",
|
||||
" print(\"Accuracy:\", accuracy_score(y, preds))\n",
|
||||
" print(\"F1:\", f1_score(y, preds))\n",
|
||||
" print(\"AUC:\", roc_auc_score(y, preds))\n",
|
||||
" print(\"Confusion:\\n\", confusion_matrix(y, preds))\n",
|
||||
" print(classification_report(y, preds))\n",
|
||||
"\n",
|
||||
" # Confusion Matrix plotten\n",
|
||||
" def plot_confusion_matrix(true_labels, predictions, label_names):\n",
|
||||
" for normalize in [None, 'true']:\n",
|
||||
" cm = confusion_matrix(true_labels, predictions, normalize=normalize)\n",
|
||||
" cm_disp = ConfusionMatrixDisplay(cm, display_labels=label_names)\n",
|
||||
" cm_disp.plot(cmap=\"Blues\")\n",
|
||||
" #cm = confusion_matrix(y, preds)\n",
|
||||
" plot_confusion_matrix(y,preds, label_names=['Low','High'])\n",
|
||||
" # plt.figure(figsize=(5,4))\n",
|
||||
" # sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", cbar=False,\n",
|
||||
" # xticklabels=[\"Predicted low\", \"Predicted high\"],\n",
|
||||
" # yticklabels=[\"Actual low\", \"Actual high\"])\n",
|
||||
" # plt.title(f\"Confusion Matrix - {title}\")\n",
|
||||
" # plt.ylabel(\"True label\")\n",
|
||||
" # plt.xlabel(\"Predicted label\")\n",
|
||||
" # plt.show()\n",
|
||||
"\n",
|
||||
"# Aufrufen für Train/Val/Test\n",
|
||||
"print(\"TRAIN:\")\n",
|
||||
"evaluate(model, X_train, y_train, title=\"Train\")\n",
|
||||
"\n",
|
||||
"print(\"VAL:\")\n",
|
||||
"evaluate(model, X_val, y_val, title=\"Validation\")\n",
|
||||
"\n",
|
||||
"print(\"TEST:\")\n",
|
||||
"evaluate(model, X_test, y_test, title=\"Test\")\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "c43b0c80",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"joblib.dump(model, \"xgb_model_with_MAD.joblib\")\n",
|
||||
"joblib.dump(normalizer, \"normalizer_with_MAD.joblib\")\n",
|
||||
"print(\"Model gespeichert.\")\n",
|
||||
"\n",
|
||||
"model.save_model(\"xgb_model_with_MAD.json\") # als JSON (lesbar, portabel)\n",
|
||||
"model.save_model(\"xgb_model_with_MAD.bin\") # als Binärdatei (kompakt)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "3195cc84",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import os\n",
|
||||
"os.getcwd()"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "Python 3 (ipykernel)",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.12.10"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,110 @@
|
||||
database:
|
||||
path: "/home/edgekit/MSY_FS/databases/database.sqlite"
|
||||
table: feature_table
|
||||
key: _Id
|
||||
|
||||
model:
|
||||
path: "/home/edgekit/MSY_FS/fahrsimulator_msy2526_ai/predict_pipeline/cnn_crossVal_EarlyFusion_V2_0103.keras"
|
||||
|
||||
scaler:
|
||||
use_scaling: true
|
||||
path: "/home/edgekit/MSY_FS/fahrsimulator_msy2526_ai/predict_pipeline/scaler_crossVal_EarlyFusion_V2_0103.joblib"
|
||||
|
||||
mqtt:
|
||||
enabled: true
|
||||
host: "141.75.223.13"
|
||||
port: 1883
|
||||
topic: "PREDICTION"
|
||||
client_id: "jetson-board"
|
||||
qos: 0
|
||||
retain: false
|
||||
# username: ""
|
||||
# password: ""
|
||||
tls:
|
||||
enabled: false
|
||||
# ca_cert: ""
|
||||
# client_cert: ""
|
||||
# client_key: ""
|
||||
publish_format:
|
||||
result_key: prediction # where to store the predicted value in payload
|
||||
include_metadata: true # e.g., timestamps, rowid, etc.
|
||||
|
||||
sample:
|
||||
columns:
|
||||
- _Id
|
||||
- start_time
|
||||
- FACE_AU01_mean
|
||||
- FACE_AU02_mean
|
||||
- FACE_AU04_mean
|
||||
- FACE_AU05_mean
|
||||
- FACE_AU06_mean
|
||||
- FACE_AU07_mean
|
||||
- FACE_AU09_mean
|
||||
- FACE_AU10_mean
|
||||
- FACE_AU11_mean
|
||||
- FACE_AU12_mean
|
||||
- FACE_AU14_mean
|
||||
- FACE_AU15_mean
|
||||
- FACE_AU17_mean
|
||||
- FACE_AU20_mean
|
||||
- FACE_AU23_mean
|
||||
- FACE_AU24_mean
|
||||
- FACE_AU25_mean
|
||||
- FACE_AU26_mean
|
||||
- FACE_AU28_mean
|
||||
- FACE_AU43_mean
|
||||
- Fix_count_short_66_150
|
||||
- Fix_count_medium_300_500
|
||||
- Fix_count_long_gt_1000
|
||||
- Fix_count_100
|
||||
- Fix_mean_duration
|
||||
- Fix_median_duration
|
||||
- Sac_count
|
||||
- Sac_mean_amp
|
||||
- Sac_mean_dur
|
||||
- Sac_median_dur
|
||||
- Blink_count
|
||||
- Blink_mean_dur
|
||||
- Blink_median_dur
|
||||
- Pupil_mean
|
||||
- Pupil_IPA
|
||||
|
||||
fill_nan_with_median: true
|
||||
discard_if_all_nan: true
|
||||
|
||||
fallback:
|
||||
FACE_AU01_mean: 0.7645230925040001
|
||||
FACE_AU02_mean: 0.731433810144
|
||||
FACE_AU04_mean: 0.19544571909800001
|
||||
FACE_AU05_mean: 0.5459417841199999
|
||||
FACE_AU06_mean: 0.11525241050400001
|
||||
FACE_AU07_mean: 0.012
|
||||
FACE_AU09_mean: 0.1025071305288
|
||||
FACE_AU10_mean: 0.018860388261559197
|
||||
FACE_AU11_mean: 0.4
|
||||
FACE_AU12_mean: 0.06147405784940001
|
||||
FACE_AU14_mean: 0.3035830324256
|
||||
FACE_AU15_mean: 0.429531116458
|
||||
FACE_AU17_mean: 0.59837751402
|
||||
FACE_AU20_mean: 0.0
|
||||
FACE_AU23_mean: 0.36847432157119997
|
||||
FACE_AU24_mean: 0.460720551004
|
||||
FACE_AU25_mean: 0.08549070376580001
|
||||
FACE_AU26_mean: 0.15669224557279998
|
||||
FACE_AU28_mean: 0.4071362423348
|
||||
FACE_AU43_mean: 0.10835549767080001
|
||||
Fix_count_short_66_150: 1.0
|
||||
Fix_count_medium_300_500: 0.0
|
||||
Fix_count_long_gt_1000: 0.0
|
||||
Fix_count_100: 1.0
|
||||
Fix_mean_duration: 60.869565217391305
|
||||
Fix_median_duration: 40.0
|
||||
Sac_count: 98.0
|
||||
Sac_mean_amp: 0.12010199338955968
|
||||
Sac_mean_dur: 263.5705727294512
|
||||
Sac_median_dur: 160.0
|
||||
Blink_count: 14.0
|
||||
Blink_mean_dur: 0.38857142857142857
|
||||
Blink_median_dur: 0.2
|
||||
Pupil_mean: 3.2823675201416016
|
||||
Pupil_IPA: 0.0036347377340156025
|
||||
@@ -0,0 +1,253 @@
|
||||
{
|
||||
"cells": [
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "fb68b447",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"## Database creation and filling (for live system) "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "0d70a13f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"import sys\n",
|
||||
"sys.path.append('/home/edgekit/MSY_FS/fahrsimulator_msy2526_ai/tools')\n",
|
||||
"import pandas as pd\n",
|
||||
"from pathlib import Path\n",
|
||||
"import db_helpers"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "ce696366",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"# TODO: set paths and table name\n",
|
||||
"database_path = Path(r\"database.sqlite\") # this path references an empty, but already created sqlite file\n",
|
||||
"parquet_path = Path(r\"...parquet\") # this path leads to the data that should be used to fill the databse\n",
|
||||
"table_name = \"XXX\" # name of the new table"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b1aa9398",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dataset = pd.read_parquet(parquet_path)\n",
|
||||
"dataset.head()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b183746e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"dataset.dtypes"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "24ed769d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"con, cursor = db_helpers.connect_db(database_path)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "7007c68f",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Select a subset to insert into database "
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e604ed30",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"df_clean = dataset.drop(columns=['subjectID','rowID', 'STUDY', 'LEVEL', 'PHASE'])\n",
|
||||
"df_first_100 = df_clean.head(200)\n",
|
||||
"df_first_100 = df_first_100.reset_index(drop=True)\n",
|
||||
"df_first_100.insert(0, '_Id', df_first_100.index + 1)"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "92171186",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Type conversion"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "e77a812e",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"def pandas_to_sqlite_dtype(dtype):\n",
|
||||
" if pd.api.types.is_integer_dtype(dtype):\n",
|
||||
" return \"INTEGER\"\n",
|
||||
" if pd.api.types.is_float_dtype(dtype):\n",
|
||||
" return \"REAL\"\n",
|
||||
" if pd.api.types.is_bool_dtype(dtype):\n",
|
||||
" return \"INTEGER\"\n",
|
||||
" if pd.api.types.is_datetime64_any_dtype(dtype):\n",
|
||||
" return \"TEXT\"\n",
|
||||
" return \"TEXT\"\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "45af9956",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Define constraints and primary key"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "0e8897b2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"columns = {\n",
|
||||
" col: pandas_to_sqlite_dtype(dtype)\n",
|
||||
" for col, dtype in df_first_100.dtypes.items()\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"constraints = {\n",
|
||||
" \"_Id\": [\"NOT NULL\"]\n",
|
||||
"}\n",
|
||||
"\n",
|
||||
"primary_key = {\n",
|
||||
" \"pk_df_first_100\": [\"_Id\"]\n",
|
||||
"}\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "markdown",
|
||||
"id": "133e92ee",
|
||||
"metadata": {},
|
||||
"source": [
|
||||
"Create the table"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "4ab57624",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"sql = db_helpers.create_table(\n",
|
||||
" conn=con,\n",
|
||||
" cursor=cursor,\n",
|
||||
" table_name=table_name,\n",
|
||||
" columns=columns,\n",
|
||||
" constraints=constraints,\n",
|
||||
" primary_key=primary_key,\n",
|
||||
" commit=True\n",
|
||||
")\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "25096a7f",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"columns_to_insert = {\n",
|
||||
" col: df_first_100[col].tolist()\n",
|
||||
" for col in df_first_100.columns\n",
|
||||
"}"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "7a5a3aa8",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"db_helpers.insert_rows_into_table(\n",
|
||||
" conn=con,\n",
|
||||
" cursor=cursor,\n",
|
||||
" table_name=table_name,\n",
|
||||
" columns=columns_to_insert,\n",
|
||||
" commit=True\n",
|
||||
")\n"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "b56beae2",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"request = db_helpers.get_data_from_table(conn=con, table_name='rawdata',columns_list=['*'])"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "a4a74a9d",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"request.head()"
|
||||
]
|
||||
},
|
||||
{
|
||||
"cell_type": "code",
|
||||
"execution_count": null,
|
||||
"id": "da0f8737",
|
||||
"metadata": {},
|
||||
"outputs": [],
|
||||
"source": [
|
||||
"db_helpers.disconnect_db(con, cursor)"
|
||||
]
|
||||
}
|
||||
],
|
||||
"metadata": {
|
||||
"kernelspec": {
|
||||
"display_name": "310",
|
||||
"language": "python",
|
||||
"name": "python3"
|
||||
},
|
||||
"language_info": {
|
||||
"codemirror_mode": {
|
||||
"name": "ipython",
|
||||
"version": 3
|
||||
},
|
||||
"file_extension": ".py",
|
||||
"mimetype": "text/x-python",
|
||||
"name": "python",
|
||||
"nbconvert_exporter": "python",
|
||||
"pygments_lexer": "ipython3",
|
||||
"version": "3.10.19"
|
||||
}
|
||||
},
|
||||
"nbformat": 4,
|
||||
"nbformat_minor": 5
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
[Unit]
|
||||
Description=Predict latest sample and send message
|
||||
After=network.target
|
||||
StartLimitIntervalSec=0
|
||||
[Service]
|
||||
Type=oneshot
|
||||
User=edgekit
|
||||
ExecStart=/home/edgekit/anaconda3/envs/p310_FS_TF/bin/python /home/edgekit/MSY_FS/fahrsimulator_msy2526_ai/predict_pipeline/predict_sample.py
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
@@ -0,0 +1,12 @@
|
||||
[Unit]
|
||||
Description=Run predict sample every 5 seconds
|
||||
|
||||
[Timer]
|
||||
OnActiveSec=60
|
||||
OnUnitActiveSec=5
|
||||
AccuracySec=1s
|
||||
Unit=predict.service
|
||||
|
||||
[Install]
|
||||
WantedBy=timers.target
|
||||
|
||||
@@ -0,0 +1,247 @@
|
||||
# Imports
|
||||
import pandas as pd
|
||||
import json
|
||||
from pathlib import Path
|
||||
import numpy as np
|
||||
import sys
|
||||
import yaml
|
||||
import pickle
|
||||
sys.path.append('/home/edgekit/MSY_FS/fahrsimulator_msy2526_ai/tools')
|
||||
import db_helpers
|
||||
import joblib
|
||||
import paho.mqtt.client as mqtt
|
||||
|
||||
def _load_serialized(path: Path):
|
||||
suffix = path.suffix.lower()
|
||||
if suffix == ".pkl":
|
||||
with path.open("rb") as f:
|
||||
return pickle.load(f)
|
||||
if suffix == ".joblib":
|
||||
return joblib.load(path)
|
||||
raise ValueError(f"Unsupported file format: {suffix}. Use .pkl or .joblib.")
|
||||
|
||||
def getLastEntryFromSQLite(path, table_name, key="_Id"):
|
||||
conn, cursor = db_helpers.connect_db(path)
|
||||
try:
|
||||
row_df = db_helpers.get_data_from_table(
|
||||
conn=conn,
|
||||
table_name=table_name,
|
||||
order_by={key: "DESC"},
|
||||
limit=1,
|
||||
)
|
||||
finally:
|
||||
db_helpers.disconnect_db(conn, cursor, commit=False)
|
||||
|
||||
if row_df.empty:
|
||||
return pd.Series(dtype="object")
|
||||
|
||||
return row_df.iloc[0]
|
||||
|
||||
def callModel(sample, model_path):
|
||||
if callable(sample):
|
||||
raise TypeError(
|
||||
f"Invalid sample type: got callable `{getattr(sample, '__name__', type(sample).__name__)}`. "
|
||||
"Expected numpy array / pandas row."
|
||||
)
|
||||
|
||||
model_path = Path(model_path)
|
||||
if not model_path.is_absolute():
|
||||
model_path = Path.cwd() / model_path
|
||||
model_path = model_path.resolve()
|
||||
|
||||
suffix = model_path.suffix.lower()
|
||||
if suffix in {".pkl", ".joblib"}:
|
||||
model = _load_serialized(model_path)
|
||||
elif suffix == ".keras":
|
||||
import tensorflow
|
||||
tensorflow.get_logger().setLevel("ERROR")
|
||||
model = tensorflow.keras.models.load_model(model_path)
|
||||
else:
|
||||
raise ValueError(f"Unsupported model format: {suffix}. Use .pkl, .joblib, or .keras.")
|
||||
|
||||
x = np.asarray(sample, dtype=np.float32)
|
||||
if x.ndim == 1:
|
||||
x = x.reshape(1, -1)
|
||||
|
||||
if suffix == ".keras":
|
||||
x_full = x
|
||||
prediction = (model.predict(x_full[:, :35], verbose=0) > 0.5).astype(int)
|
||||
|
||||
else:
|
||||
if hasattr(model, "predict"):
|
||||
prediction = model.predict(x[:,:20])
|
||||
elif callable(model):
|
||||
prediction = model(x[:,:20])
|
||||
else:
|
||||
raise TypeError("Loaded model has no .predict(...) and is not callable.")
|
||||
|
||||
prediction = np.asarray(prediction)
|
||||
if prediction.size == 1:
|
||||
return prediction.item()
|
||||
return prediction.squeeze()
|
||||
|
||||
def buildMessage(valid, result: np.int32, config_file_path, sample=None):
|
||||
with Path(config_file_path).open("r", encoding="utf-8") as f:
|
||||
cfg = yaml.safe_load(f)
|
||||
|
||||
mqtt_cfg = cfg.get("mqtt", {})
|
||||
result_key = mqtt_cfg.get("publish_format", {}).get("result_key", "prediction")
|
||||
|
||||
sample_id = None
|
||||
if isinstance(sample, pd.Series):
|
||||
sample_id = sample.get("_Id", sample.get("_id"))
|
||||
elif isinstance(sample, dict):
|
||||
sample_id = sample.get("_Id", sample.get("_id"))
|
||||
|
||||
message = {
|
||||
"valid": bool(valid),
|
||||
"_id": sample_id,
|
||||
result_key: np.asarray(result).tolist() if isinstance(result, np.ndarray) else result,
|
||||
}
|
||||
return message
|
||||
|
||||
def convert_int64(obj):
|
||||
if isinstance(obj, np.int64):
|
||||
return int(obj)
|
||||
# If the object is a dictionary or list, recursively convert its values
|
||||
elif isinstance(obj, dict):
|
||||
return {key: convert_int64(value) for key, value in obj.items()}
|
||||
elif isinstance(obj, list):
|
||||
return [convert_int64(item) for item in obj]
|
||||
return obj
|
||||
|
||||
def sendMessage(config_file_path, message):
|
||||
# Load the configuration
|
||||
with Path(config_file_path).open("r", encoding="utf-8") as f:
|
||||
cfg = yaml.safe_load(f)
|
||||
|
||||
# Get MQTT configuration
|
||||
mqtt_cfg = cfg.get("mqtt", {})
|
||||
topic = mqtt_cfg.get("topic", "ml/predictions")
|
||||
|
||||
# Convert message to ensure no np.int64 values remain
|
||||
message = convert_int64(message)
|
||||
|
||||
# Serialize the message to JSON
|
||||
payload = json.dumps(message, ensure_ascii=False)
|
||||
print(payload)
|
||||
|
||||
# publish via MQTT using config parameters above.
|
||||
client = mqtt.Client(client_id=mqtt_cfg.get("client_id", "predictor-01"),
|
||||
callback_api_version=mqtt.CallbackAPIVersion.VERSION2)
|
||||
callback_api_version=mqtt.CallbackAPIVersion.VERSION2
|
||||
if "username" in mqtt_cfg and mqtt_cfg.get("username"):
|
||||
client.username_pw_set(mqtt_cfg["username"], mqtt_cfg.get("password"))
|
||||
client.connect(mqtt_cfg.get("host", "localhost"), int(mqtt_cfg.get("port", 1883)), 60)
|
||||
client.publish(
|
||||
topic=topic,
|
||||
payload=payload,
|
||||
qos=int(mqtt_cfg.get("qos", 1)),
|
||||
retain=bool(mqtt_cfg.get("retain", False)),
|
||||
)
|
||||
client.disconnect()
|
||||
return
|
||||
|
||||
def replace_nan(sample, config_file_path: Path):
|
||||
with config_file_path.open("r", encoding="utf-8") as f:
|
||||
cfg = yaml.safe_load(f)
|
||||
|
||||
fallback_map = cfg.get("fallback", {})
|
||||
|
||||
if sample.empty:
|
||||
return False, sample
|
||||
|
||||
nan_ratio = sample.isna().mean()
|
||||
valid = nan_ratio <= 0.5
|
||||
|
||||
if valid and fallback_map:
|
||||
sample = sample.fillna(value=fallback_map)
|
||||
|
||||
return valid, sample
|
||||
|
||||
def sample_to_numpy(sample, drop_cols=("_Id", "start_time")):
|
||||
if isinstance(sample, pd.Series):
|
||||
sample = sample.drop(labels=list(drop_cols), errors="ignore")
|
||||
return sample.to_numpy()
|
||||
|
||||
if isinstance(sample, pd.DataFrame):
|
||||
sample = sample.drop(columns=list(drop_cols), errors="ignore")
|
||||
return sample.to_numpy()
|
||||
|
||||
return np.asarray(sample)
|
||||
|
||||
def scale_sample(sample, use_scaling=False, scaler_path=None):
|
||||
if not use_scaling or scaler_path is None:
|
||||
return sample
|
||||
scaler_path = Path(scaler_path)
|
||||
if not scaler_path.is_absolute():
|
||||
scaler_path = Path.cwd() / scaler_path
|
||||
scaler_path = scaler_path.resolve()
|
||||
normalizer = _load_serialized(scaler_path)
|
||||
|
||||
# normalizer format from model_training/tools/scaler.py:
|
||||
# {"scalers": {...}, "method": "...", "scope": "..."}
|
||||
scalers = normalizer.get("scalers", {}) if isinstance(normalizer, dict) else {}
|
||||
scope = normalizer.get("scope", "global") if isinstance(normalizer, dict) else "global"
|
||||
if scope == "global":
|
||||
scaler = scalers.get("global")
|
||||
else:
|
||||
scaler = scalers.get("global", next(iter(scalers.values()), None))
|
||||
|
||||
# Optional fallback if the stored object is already a raw scaler.
|
||||
if scaler is None and hasattr(normalizer, "transform"):
|
||||
scaler = normalizer
|
||||
if scaler is None or not hasattr(scaler, "transform"):
|
||||
return sample
|
||||
|
||||
df = sample.to_frame().T if isinstance(sample, pd.Series) else sample.copy()
|
||||
feature_names = getattr(scaler, "feature_names_in_", None)
|
||||
if feature_names is None:
|
||||
return sample
|
||||
|
||||
# Keep columns not in the normalizer unchanged.
|
||||
cols_to_scale = [c for c in df.columns if c in set(feature_names)]
|
||||
if cols_to_scale:
|
||||
df.loc[:, cols_to_scale] = scaler.transform(df.loc[:, cols_to_scale])
|
||||
|
||||
return df.iloc[0] if isinstance(sample, pd.Series) else df
|
||||
|
||||
def main():
|
||||
pd.set_option('future.no_silent_downcasting', True) # kann ggf raus
|
||||
|
||||
config_file_path = Path("/home/edgekit/MSY_FS/fahrsimulator_msy2526_ai/predict_pipeline/config.yaml")
|
||||
with config_file_path.open("r", encoding="utf-8") as f:
|
||||
cfg = yaml.safe_load(f)
|
||||
|
||||
database_path = cfg["database"]["path"]
|
||||
table_name = cfg["database"]["table"]
|
||||
row_key = cfg["database"]["key"]
|
||||
|
||||
|
||||
sample = getLastEntryFromSQLite(database_path, table_name, row_key)
|
||||
valid, sample = replace_nan(sample, config_file_path=config_file_path)
|
||||
|
||||
if not valid:
|
||||
print("Sample invalid: more than 50% NaN.")
|
||||
message = buildMessage(valid, None, config_file_path, sample=sample)
|
||||
sendMessage(config_file_path, message)
|
||||
return
|
||||
|
||||
model_path = cfg["model"]["path"]
|
||||
scaler_path = cfg["scaler"]["path"]
|
||||
use_scaling = cfg["scaler"]["use_scaling"]
|
||||
|
||||
sample = scale_sample(sample, use_scaling=use_scaling, scaler_path=scaler_path)
|
||||
sample_np = sample_to_numpy(sample)
|
||||
|
||||
prediction = callModel(model_path=model_path, sample=sample_np)
|
||||
|
||||
message = buildMessage(valid, prediction, config_file_path, sample=sample)
|
||||
sendMessage(config_file_path, message)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,216 @@
|
||||
# Predict Service and Timer Documentation
|
||||
|
||||
## Overview
|
||||
|
||||
This setup uses **systemd services and timers** to repeatedly execute a
|
||||
Python script that performs prediction on the latest sample and sends a
|
||||
message.
|
||||
|
||||
The systemd unit files are typically stored in:
|
||||
|
||||
/etc/systemd/system/
|
||||
|
||||
For this setup, the relevant files are:
|
||||
|
||||
/etc/systemd/system/predict.service
|
||||
/etc/systemd/system/predict.timer
|
||||
|
||||
These files define the service execution and the timer scheduling.
|
||||
|
||||
- `predict.service` -- defines how the Python script is executed
|
||||
- `predict.timer` -- schedules the repeated execution of the service
|
||||
|
||||
The timer triggers the service **every 5 seconds** after the first
|
||||
activation.
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
# Systemd Timer
|
||||
|
||||
File: `predict.timer`
|
||||
|
||||
``` ini
|
||||
[Unit]
|
||||
Description=Run predict sample every 5 seconds
|
||||
|
||||
[Timer]
|
||||
OnActiveSec=60
|
||||
OnUnitActiveSec=5
|
||||
AccuracySec=1s
|
||||
Unit=predict.service
|
||||
|
||||
[Install]
|
||||
WantedBy=timers.target
|
||||
```
|
||||
|
||||
## Behavior
|
||||
|
||||
- **OnActiveSec=60**\
|
||||
The timer starts **60 seconds after it is activated**.
|
||||
|
||||
- **OnUnitActiveSec=5**\
|
||||
After the service has run once, it will be triggered again **every 5
|
||||
seconds**.
|
||||
|
||||
- **AccuracySec=1s**\
|
||||
Allows systemd to schedule the timer with **1 second precision**.
|
||||
|
||||
- **Unit=predict.service**\
|
||||
Defines which service should be triggered by the timer.
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
# Systemd Service
|
||||
|
||||
File: `predict.service`
|
||||
|
||||
``` ini
|
||||
[Unit]
|
||||
Description=Predict latest sample and send message
|
||||
After=network.target
|
||||
StartLimitIntervalSec=0
|
||||
|
||||
[Service]
|
||||
Type=oneshot
|
||||
User=edgekit
|
||||
ExecStart=/home/edgekit/anaconda3/envs/p310_FS_TF/bin/python /home/edgekit/MSY_FS/fahrsimulator_msy2526_ai/predict_pipeline/predict_sample.py
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
```
|
||||
|
||||
## Behavior
|
||||
|
||||
- **Type=oneshot**\
|
||||
The service runs the script once and then exits.
|
||||
|
||||
- **User=edgekit**\
|
||||
The script is executed under the `edgekit` user.
|
||||
|
||||
- **ExecStart**\
|
||||
Executes the Python script using the specified conda environment.
|
||||
|
||||
- **After=network.target**\
|
||||
Ensures the service only runs after the network is available.
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
# Execution Flow
|
||||
|
||||
1. The **timer starts** after it is enabled.
|
||||
2. After **60 seconds**, the first execution happens (this results from the duration of the camera processing initialization)
|
||||
3. The timer triggers `predict.service`.
|
||||
4. The service runs `predict_sample.py`.
|
||||
5. Once the script finishes, the service exits.
|
||||
6. The timer triggers the service again **every 5 seconds**.
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
# Debugging and Monitoring
|
||||
|
||||
## View Live Output
|
||||
|
||||
All `print()` output from the Python script is written to the **systemd
|
||||
journal**.
|
||||
|
||||
Follow the output live with:
|
||||
|
||||
``` bash
|
||||
journalctl -u predict.service -f
|
||||
```
|
||||
|
||||
This command is typically the most useful for debugging.
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
# Common Systemd Commands
|
||||
|
||||
## Check Service Status
|
||||
|
||||
``` bash
|
||||
systemctl status predict.service
|
||||
```
|
||||
|
||||
Shows the last execution result and recent log lines.
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
## Check Timer Status
|
||||
|
||||
``` bash
|
||||
systemctl status predict.timer
|
||||
```
|
||||
|
||||
Shows when the timer last ran and when it will run next.
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
## List All Timers
|
||||
|
||||
``` bash
|
||||
systemctl list-timers
|
||||
```
|
||||
|
||||
Displays all active timers and their next scheduled execution.
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
# Manual Execution
|
||||
|
||||
To run the service manually once:
|
||||
|
||||
``` bash
|
||||
systemctl start predict.service
|
||||
```
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
# Restarting the Systemd Units
|
||||
|
||||
## Restart the Service
|
||||
|
||||
``` bash
|
||||
systemctl restart predict.service
|
||||
```
|
||||
|
||||
## Restart the Timer
|
||||
|
||||
``` bash
|
||||
systemctl restart predict.timer
|
||||
```
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
## Reload Systemd After Changes
|
||||
|
||||
If `.service` or `.timer` files were modified:
|
||||
|
||||
``` bash
|
||||
systemctl daemon-reload
|
||||
systemctl restart predict.timer
|
||||
systemctl restart predict.service
|
||||
```
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
## Enabling the Timer
|
||||
|
||||
To ensure the timer starts automatically on system boot:
|
||||
|
||||
``` bash
|
||||
systemctl enable predict.timer
|
||||
systemctl start predict.timer
|
||||
```
|
||||
|
||||
------------------------------------------------------------------------
|
||||
|
||||
# Summary
|
||||
|
||||
- `predict.timer` schedules periodic execution.
|
||||
- `predict.service` runs the Python prediction script.
|
||||
- The script runs **every 5 seconds** after the initial delay.
|
||||
- Logs and script output are available through:
|
||||
|
||||
``` bash
|
||||
journalctl -u predict.service -f
|
||||
```
|
||||
@@ -0,0 +1,196 @@
|
||||
name: 'prediction_env'
|
||||
channels:
|
||||
- defaults
|
||||
- conda-forge
|
||||
dependencies:
|
||||
- _py-xgboost-mutex=2.0=cpu_2
|
||||
- absl-py=2.3.1=py310haa95532_0
|
||||
- aom=3.12.1=h00a0c3c_0
|
||||
- arrow-cpp=21.0.0=hcdc3a1c_2
|
||||
- asttokens=3.0.1=pyhd8ed1ab_0
|
||||
- astunparse=1.6.3=py_0
|
||||
- aws-c-auth=0.9.0=h02ab6af_2
|
||||
- aws-c-cal=0.9.2=h02ab6af_1
|
||||
- aws-c-common=0.12.4=h02ab6af_0
|
||||
- aws-c-compression=0.3.1=h02ab6af_2
|
||||
- aws-c-event-stream=0.5.6=h02ab6af_0
|
||||
- aws-c-http=0.10.4=h02ab6af_0
|
||||
- aws-c-io=0.21.4=h02ab6af_0
|
||||
- aws-c-mqtt=0.13.3=h02ab6af_0
|
||||
- aws-c-s3=0.8.7=h02ab6af_0
|
||||
- aws-c-sdkutils=0.2.4=h02ab6af_1
|
||||
- aws-checksums=0.2.7=h02ab6af_1
|
||||
- aws-crt-cpp=0.34.0=h885b0b7_0
|
||||
- aws-sdk-cpp=1.11.638=hf0af688_0
|
||||
- blas=1.0=mkl
|
||||
- brotlicffi=1.2.0.0=py310h885b0b7_0
|
||||
- bzip2=1.0.8=h2bbff1b_6
|
||||
- c-ares=1.34.6=h2c209ce_0
|
||||
- ca-certificates=2026.1.4=h4c7d964_0
|
||||
- cairo=1.18.4=he9e932c_0
|
||||
- certifi=2026.01.04=py310haa95532_0
|
||||
- cffi=2.0.0=py310h02ab6af_1
|
||||
- charset-normalizer=3.4.4=py310haa95532_0
|
||||
- colorama=0.4.6=pyhd8ed1ab_1
|
||||
- comm=0.2.3=pyhe01879c_0
|
||||
- dav1d=1.2.1=h2bbff1b_0
|
||||
- debugpy=1.8.20=py310h699e580_0
|
||||
- decorator=5.2.1=pyhd8ed1ab_0
|
||||
- exceptiongroup=1.3.1=pyhd8ed1ab_0
|
||||
- executing=2.2.1=pyhd8ed1ab_0
|
||||
- expat=2.7.4=hd7fb8db_0
|
||||
- flatbuffers=24.3.25=h21716d4_0
|
||||
- fontconfig=2.15.0=hd211d86_0
|
||||
- freeglut=3.8.0=hfcef157_0
|
||||
- freetype=2.14.1=hfbffc0b_0
|
||||
- fribidi=1.0.16=haf45083_0
|
||||
- gast=0.7.0=pyhd3eb1b0_0
|
||||
- gflags=2.2.2=hd77b12b_1
|
||||
- giflib=5.2.2=h7edc060_0
|
||||
- glog=0.5.0=hd77b12b_1
|
||||
- google-pasta=0.2.0=pyhd3eb1b0_0
|
||||
- graphite2=1.3.14=hd77b12b_1
|
||||
- grpcio=1.74.1=py310h5c751cc_0
|
||||
- h5py=3.15.1=py310he283ef2_1
|
||||
- harfbuzz=12.3.0=h3ef6528_1
|
||||
- hdf5=1.14.5=ha36df97_2
|
||||
- icc_rt=2022.1.0=h6049295_2
|
||||
- icu=73.1=h6c2663c_0
|
||||
- idna=3.11=py310haa95532_0
|
||||
- intel-openmp=2025.0.0=haa95532_1164
|
||||
- ipykernel=7.2.0=pyh6dadd2b_1
|
||||
- ipython=8.37.0=pyha7b4d00_0
|
||||
- jedi=0.19.2=pyhd8ed1ab_1
|
||||
- joblib=1.5.3=py310haa95532_0
|
||||
- jpeg=9f=ha349fce_0
|
||||
- jupyter_client=8.8.0=pyhcf101f3_0
|
||||
- jupyter_core=5.9.1=pyh6dadd2b_0
|
||||
- keras=3.11.2=py310h51baaa3_0
|
||||
- krb5=1.21.3=hdf4eb48_0
|
||||
- lcms2=2.17=h3732fa5_0
|
||||
- lerc=4.0.0=h5da7b33_0
|
||||
- libabseil=20250814.1=cxx17_hcd311fc_0
|
||||
- libavif=1.3.0=h5bd13ec_0
|
||||
- libbrotlicommon=1.2.0=h907acca_0
|
||||
- libbrotlidec=1.2.0=h02c67a5_0
|
||||
- libbrotlienc=1.2.0=h483e6b9_0
|
||||
- libcurl=8.17.0=h6e672f4_1
|
||||
- libdeflate=1.22=h5bf469e_0
|
||||
- libexpat=2.7.4=hd7fb8db_0
|
||||
- libffi=3.4.4=hd77b12b_1
|
||||
- libglib=2.86.3=h9bccc14_0
|
||||
- libgrpc=1.74.1=hde67744_0
|
||||
- libhwloc=2.12.1=default_hfa10c62_1000
|
||||
- libiconv=1.16=h2bbff1b_3
|
||||
- libkrb5=1.22.1=hb237eb7_0
|
||||
- libopenjpeg=2.5.4=h02ab6af_1
|
||||
- libpng=1.6.54=ha15c746_0
|
||||
- libprotobuf=6.33.0=h2a56892_1
|
||||
- libre2-11=2025.11.05=ha6b10e7_0
|
||||
- libsodium=1.0.20=hc70643c_0
|
||||
- libssh2=1.11.1=h2addb87_0
|
||||
- libthrift=0.22.0=ha2884a9_0
|
||||
- libtiff=4.7.1=h3a18249_0
|
||||
- libwebp-base=1.6.0=hbf3958f_0
|
||||
- libxgboost=3.1.2=h585ebfc_0
|
||||
- libxml2=2.13.9=h6201b9f_0
|
||||
- libzlib=1.3.1=h02ab6af_0
|
||||
- lz4-c=1.9.4=h2bbff1b_1
|
||||
- m2w64-gcc-libgfortran=5.3.0=6
|
||||
- m2w64-gcc-libs=5.3.0=7
|
||||
- m2w64-gcc-libs-core=5.3.0=7
|
||||
- m2w64-gmp=6.1.0=2
|
||||
- m2w64-libwinpthread-git=5.0.0.4634.697f757=2
|
||||
- markdown=3.10=py310haa95532_0
|
||||
- markdown-it-py=2.2.0=py310haa95532_1
|
||||
- markupsafe=3.0.2=py310h827c3e9_0
|
||||
- matplotlib-inline=0.2.1=pyhd8ed1ab_0
|
||||
- mdurl=0.1.2=py310haa95532_0
|
||||
- mkl=2025.0.0=h5da7b33_930
|
||||
- mkl-service=2.5.2=py310h0b37514_0
|
||||
- mkl_fft=2.1.1=py310h300f80d_0
|
||||
- mkl_random=1.3.0=py310ha5e6156_0
|
||||
- ml_dtypes=0.5.4=py310h42c1672_0
|
||||
- mpi=1.0=msmpi
|
||||
- mpi4py=4.0.3=py310h02ab6af_1
|
||||
- msmpi=10.1.1=had4844c_0
|
||||
- msys2-conda-epoch=20160418=1
|
||||
- namex=0.1.0=py310haa95532_0
|
||||
- nest-asyncio=1.6.0=pyhd8ed1ab_1
|
||||
- numpy-base=2.1.3=py310he4e2855_3
|
||||
- openssl=3.6.1=hf411b9b_1
|
||||
- opt_einsum=3.3.0=pyhd3eb1b0_1
|
||||
- optree=0.18.0=py310h03f52e7_0
|
||||
- orc=2.2.0=h79e1e1e_1
|
||||
- packaging=25.0=py310haa95532_1
|
||||
- paho-mqtt=2.1.0=pyhe01879c_1
|
||||
- parso=0.8.6=pyhcf101f3_0
|
||||
- pcre2=10.46=h5740b90_0
|
||||
- pickleshare=0.7.5=pyhd8ed1ab_1004
|
||||
- pillow=12.1.0=py310h6b7a805_0
|
||||
- pip=26.0.1=pyhc872135_0
|
||||
- pixman=0.46.4=h4043f72_0
|
||||
- platformdirs=4.9.2=pyhcf101f3_0
|
||||
- prompt-toolkit=3.0.52=pyha770c72_0
|
||||
- protobuf=6.33.0=py310ha4c6e68_0
|
||||
- psutil=7.2.2=py310h1637853_0
|
||||
- pure_eval=0.2.3=pyhd8ed1ab_1
|
||||
- py-xgboost=3.1.2=py310haa95532_0
|
||||
- pyarrow=21.0.0=py310h42c1672_1
|
||||
- pycparser=2.23=py310haa95532_0
|
||||
- pygments=2.19.2=pyhd8ed1ab_0
|
||||
- pysocks=1.7.1=py310haa95532_1
|
||||
- python=3.10.19=h981015d_0
|
||||
- python-dateutil=2.9.0.post0=pyhe01879c_2
|
||||
- python-flatbuffers=24.3.25=py310haa95532_0
|
||||
- python_abi=3.10=2_cp310
|
||||
- pywin32=311=py310h282bd7d_1
|
||||
- pyyaml=6.0.3=py310hb9a58be_0
|
||||
- pyzmq=27.1.0=py310h535538e_0
|
||||
- re2=2025.11.05=hc24cdf5_0
|
||||
- requests=2.32.5=py310haa95532_1
|
||||
- rich=14.2.0=py310haa95532_0
|
||||
- scipy=1.15.3=py310h1bbe36f_1
|
||||
- setuptools=80.10.2=py310haa95532_0
|
||||
- six=1.17.0=pyhe01879c_1
|
||||
- snappy=1.2.2=hab6b7b3_1
|
||||
- sqlite=3.51.1=hda9a48d_0
|
||||
- stack_data=0.6.3=pyhd8ed1ab_1
|
||||
- tbb=2022.3.0=h90c84d6_0
|
||||
- tbb-devel=2022.3.0=h90c84d6_0
|
||||
- tensorboard=2.20.0=py310haa95532_0
|
||||
- tensorboard-data-server=0.7.0=py310haa95532_1
|
||||
- tensorflow=2.20.0=cpu_py310h6605a60_0
|
||||
- tensorflow-base=2.20.0=cpu_py310hce87ebc_0
|
||||
- termcolor=3.2.0=py310haa95532_0
|
||||
- threadpoolctl=3.5.0=py310h4442805_1
|
||||
- tk=8.6.15=hf199647_0
|
||||
- tornado=6.5.4=py310h29418f3_0
|
||||
- traitlets=5.14.3=pyhd8ed1ab_1
|
||||
- typing-extensions=4.15.0=py310haa95532_0
|
||||
- typing_extensions=4.15.0=py310haa95532_0
|
||||
- ucrt=10.0.22621.0=haa95532_0
|
||||
- urllib3=2.6.3=py310haa95532_0
|
||||
- utf8proc=2.6.1=h2bbff1b_1
|
||||
- vc=14.42=haa95532_5
|
||||
- vc14_runtime=14.44.35208=h4927774_10
|
||||
- vs2015_runtime=14.44.35208=ha6b5a95_10
|
||||
- wcwidth=0.6.0=pyhd8ed1ab_0
|
||||
- werkzeug=3.1.3=py310haa95532_0
|
||||
- wheel=0.46.3=py310haa95532_0
|
||||
- win_inet_pton=1.1.0=py310haa95532_1
|
||||
- wrapt=2.0.1=py310h02ab6af_0
|
||||
- xgboost=3.1.2=py310haa95532_0
|
||||
- xz=5.6.4=h4754444_1
|
||||
- yaml=0.2.5=he774522_0
|
||||
- zeromq=4.3.5=h5bddc39_9
|
||||
- zlib=1.3.1=h02ab6af_0
|
||||
- zstd=1.5.7=h56299aa_0
|
||||
- pip:
|
||||
- numpy==1.24.4
|
||||
- pandas==2.3.0
|
||||
- pyocclient==0.6
|
||||
- pytz==2025.2
|
||||
- scikit-learn==1.6.1
|
||||
- tzdata==2025.3
|
||||
|
||||
@@ -0,0 +1,551 @@
|
||||
# Project Report: Multimodal Driver State Analysis
|
||||
|
||||
## 1) Project Scope
|
||||
|
||||
This repository implements an end-to-end workflow for multimodal driver-state analysis in a simulator setup.
|
||||
The system combines:
|
||||
- Facial Action Units (AUs)
|
||||
- Eye-tracking features (fixations, saccades, blinks, pupil behavior)
|
||||
|
||||
Apart from this, several machine learning model architectures are presented and evaluated.
|
||||
|
||||
Content:
|
||||
- Dataset generation
|
||||
- Exploratory data analysis
|
||||
- Model training experiments
|
||||
- Real-time inference with SQlite, systemd and MQTT
|
||||
- Repository file inventory
|
||||
- Additional nformation
|
||||
|
||||
|
||||
## 2) Dataset generation
|
||||
|
||||
### 2.1 Data Access, Filtering, and Data Conversion
|
||||
|
||||
Main scripts:
|
||||
- `dataset_creation/create_parquet_files_from_owncloud.py`
|
||||
- `dataset_creation/parquet_file_creation.py`
|
||||
|
||||
Purpose:
|
||||
- Download and/or access dataset files (either download first via ```EDA/owncloud_file_access.ipynb``` or all in one with ```dataset_creation/create_parquet_files_from_owncloud.py```
|
||||
- Keep relevant columns (FACE_AUs and eye-tracking raw values)
|
||||
- Filter invalid samples (e.g., invalid level segments): Make sure not to drop rows where NaN is necessary for later feature creation, therefore use subset argument in dropNa()!
|
||||
- Export subject-level parquet files
|
||||
- Before running the scripts: be aware that the whole dataset contains 30 files with around 900 Mbytes each, provide enough storage and expect this to take a while.
|
||||
|
||||
|
||||
### 2.2 Feature Engineering (Offline)
|
||||
|
||||
Main script:
|
||||
- `dataset_creation/combined_feature_creation.py`
|
||||
|
||||
Behavior:
|
||||
- Builds fixed-size sliding windows over subject time series (window size and step size can be adjusted)
|
||||
- Uses prepared parquet files from 2.1
|
||||
- Aggregates AU statistics per window (e.g., `FACE_AUxx_mean`)
|
||||
- Computes eye-feature aggregates (fix/sacc/blink/pupil metrics)
|
||||
- Produces training-ready feature tables = dataset
|
||||
- Parameter ```MIN_DUR_BLINKS``` can be adjusted, although this value needs to make sense in combination with your sampling frequency
|
||||
- With low videostream rates, consider to reevaluate the meaningfulness of some eye-tracking features, especially the fixations
|
||||
- running the script requires a manual installation of [pygaze Analyser library](https://github.com/esdalmaijer/PyGazeAnalyser.git) from github
|
||||
|
||||
### 2.3 Online Camera + Eye + AU Feature Extraction
|
||||
|
||||
Main scripts:
|
||||
- `dataset_creation/camera_handling/camera_stream_AU_and_ET_new.py`
|
||||
- `dataset_creation/camera_handling/eyeFeature_new.py`
|
||||
- `dataset_creation/camera_handling/db_helper.py`
|
||||
|
||||
Runtime behavior:
|
||||
- Captures webcam stream with OpenCV
|
||||
- Extracts gaze/iris-based signals via MediaPipe
|
||||
- Records overlapping windows (`VIDEO_DURATION=50s`, `START_INTERVAL=5s`, `FPS=25`)
|
||||
- Runs AU extraction (`py-feat`) from recorded video segments
|
||||
- Explanation of the py-feat functionality is located in `dataset_creation/AU_creation/pyfeat_docu.ipynb`
|
||||
- Computes eye-feature summary from generated gaze parquet
|
||||
- Writes merged rows to SQLite table `feature_table`
|
||||
|
||||
Operational note:
|
||||
- `DB_PATH` and other paths are currently code-configured and must be adapted per deployment.
|
||||
|
||||
### 2.4 Two Approaches to Eye-Tracking Data Collection
|
||||
|
||||
Eye-tracking can be implemented using two main approaches:
|
||||
|
||||
Used Approach: Relative Iris Position
|
||||
- Tracks the position of the pupil within the eye region
|
||||
- The position is normalized relative to the eye itself
|
||||
- No reference to the screen or physical environment is required
|
||||
|
||||
Not Used Approach: Screen Calibration
|
||||
- Requires the user to look at 9 predefined points on the screen
|
||||
- A mapping model is trained based on these points
|
||||
- Establishes a relationship between eye movement and screen coordinates
|
||||
|
||||
Important Considerations (for both methods)
|
||||
- Keep the head as still as possible
|
||||
- Ensure consistent and even lighting conditions
|
||||
|
||||
## 3) EDA
|
||||
The directory EDA provides several files to get insights into both the raw data from AdaBase and your own dataset.
|
||||
|
||||
- `EDA.ipynb` - Main EDA notebook: recreates the plot from AdaBase documentation, lists all experiments and in general serves as a playground for you to get to know the files.
|
||||
- `distribution_plots.ipynb` - This notebook aimes to visualize the data distributions for each experiment - the goal is the find out, whether the split of experiments into high and low cognitive load is clearer if some experiments are dropped.
|
||||
- `histogramms.ipynb` - Histogram analysis of low load vs high load per feature. Additionaly, scatter plots per feature are available.
|
||||
- `researchOnSubjectPerformance.ipynb` - This noteboooks aims to see how the performance values range for the 30 subjects. The code creates and saves a table in csv-format, which will later be used as the foundation of the performance based split in ```model_training/tools/performance_based_split```
|
||||
- `owncloud_file_access.ipynb` - Get access to the files via owncloud and safe them as .h5 files, in correspondence to the parquet file creation script
|
||||
- `login.yaml` -Used to store URL and password to access files from owncloud, used in previous notebook
|
||||
- `calculate_replacement_values.ipynb` -Fallback / median computation notebook for deployment, creation of yaml syntax embedding
|
||||
|
||||
General information:
|
||||
- Due to their size, its absolutely recommended to download and save the dataset files once in the beginning
|
||||
- For better data understanding, read the [AdaBase publication](https://www.mdpi.com/1424-8220/23/1/340)
|
||||
|
||||
|
||||
## 4) Model Training
|
||||
|
||||
Included model families:
|
||||
- CNN variants (different fusion strategies)
|
||||
- XGBoost
|
||||
- Isolation Forest*
|
||||
- OCSVM*
|
||||
- DeepSVDD*
|
||||
|
||||
\* These training strategies are unsupervised, which means only low cognitive load samples are used for training. Validation then also considers high low samples.
|
||||
|
||||
|
||||
Supporting utilities in ```model_training/tools```:
|
||||
- `scaler.py`: Functions to fit, transform, save and load either MinMaxScaler or StandardScaler, subject-wise and globally - for new subjects, a fallback scaler (using mean of all subjects scaling parameters) is used
|
||||
- `performance_split.py`: Provides a function to split a group of subjects based on their performance in the AdaBase experiments, based on the results created in `researchOnSubjectPerformance.ipynb`. To split into three groups for train, validation & test, call the function twice
|
||||
- `mad_outlier_removal.py`: Functions to fit and transform data with MAD outlier removal
|
||||
- `evaluation_tools.py`: Especially used for Isolation Forest, Functions for ROC curve as well as confusion matrix
|
||||
|
||||
|
||||
### 4.1 CNNs
|
||||
This section summarizes all CNN‑based supervised learning approaches implemented in the project.
|
||||
All models operate on facial Action Unit (AU) features and, depending on the notebook, additional eye‑tracking features.
|
||||
The notebooks differ in evaluation methodology, fusion strategy, and experimental intention.
|
||||
|
||||
### 4.1.1 Baseline CNN (Notebook: *CNN_simple*)
|
||||
The first notebook implements a simple 1D CNN to establish a baseline for AU‑only classification.
|
||||
The model uses two convolutional layers, batch normalization, max pooling, and a regularized dense head.
|
||||
A single subject‑exclusive train/validation/test split is used.
|
||||
|
||||
The intention behind this notebook is to:
|
||||
- Provide a baseline performance level
|
||||
- Validate that AU features contain discriminative information
|
||||
- Identify overfitting tendencies before moving to more rigorous evaluation
|
||||
|
||||
|
||||
### 4.1.2 Cross‑Validated CNN (Notebook: *CNN_crossVal*)
|
||||
This notebook introduces 5‑fold GroupKFold cross‑validation, ensuring subject‑exclusive folds.
|
||||
The architecture is similar to the baseline but includes stronger regularization and a lower learning rate.
|
||||
|
||||
The intention behind this notebook is to:
|
||||
- Provide robust generalization estimates
|
||||
- Reduce variance caused by single‑split evaluation
|
||||
- Establish a cross‑validated AU‑only benchmark
|
||||
|
||||
|
||||
### 4.1.3 Cross‑Validated CNN (Face AUs Only) (Notebook: *CNN_crossVal_faceAUs*)
|
||||
This notebook is a streamlined version of the previous one, removing unused eye‑tracking features and focusing exclusively on AUs.
|
||||
|
||||
The intention behind this notebook is to:
|
||||
- Provide a clean AU‑only benchmark
|
||||
- Improve reproducibility and interpretability
|
||||
- Prepare for multimodal comparisons
|
||||
|
||||
|
||||
### 4.1.4 Cross‑Validated CNN with Early Fusion (AUs + Eye Features)
|
||||
(Notebook: *CNN_crossVal_faceAUs_eyeFeatures*)
|
||||
This notebook introduces early fusion, concatenating AU and eye‑tracking features into a single input vector.
|
||||
The architecture remains identical to AU‑only models.
|
||||
|
||||
The intention behind this notebook is to:
|
||||
- Evaluate whether multimodal early fusion improves performance
|
||||
- Establish a first multimodal baseline
|
||||
- Analyze class‑specific behavior via confusion matrices
|
||||
|
||||
This notebook didn't lead to any useful results.
|
||||
|
||||
### 4.1.5 Cross‑Validated CNN with Early Fusion (Refined Version) (Notebook: *CNN_crossVal_EarlyFusion*)
|
||||
This notebook refines the early‑fusion approach by removing samples with missing values and ensuring consistent multimodal input quality.
|
||||
|
||||
The intention behind this notebook is to:
|
||||
- Provide a clean and fully validated early‑fusion model
|
||||
- Investigate multimodal complementarity under rigorous CV
|
||||
- Improve interpretability through aggregated confusion matrices
|
||||
|
||||
|
||||
### 4.1.6 Cross‑Validated CNN with Early Fusion and Subset Filtering (Notebook: *CNN_crossVal_EarlyFusion_Filter*)
|
||||
This notebook applies domain‑specific filtering to isolate a more homogeneous subset of cognitive states before training.
|
||||
|
||||
The intention behind this notebook is to:
|
||||
- Evaluate whether subset filtering improves multimodal learning
|
||||
- Reduce dataset heterogeneity
|
||||
- Provide a controlled multimodal benchmark
|
||||
|
||||
|
||||
### 4.1.7 Hybrid‑Fusion CNN (Notebook: *CNN_crossVal_HybridFusion*)
|
||||
This notebook introduces a hybrid‑fusion architecture with two modality‑specific branches:
|
||||
- A 1D CNN for AUs
|
||||
- A dense MLP for eye‑tracking features
|
||||
|
||||
The branches are fused before classification.
|
||||
|
||||
The intention behind this notebook is to:
|
||||
- Allow each modality to learn specialized representations
|
||||
- Evaluate whether hybrid fusion outperforms early fusion
|
||||
- Provide a strong multimodal benchmark
|
||||
|
||||
|
||||
### 4.1.8 Early‑Fusion CNN with Independent Test Evaluation (Notebook: *CNN_crossVal_EarlyFusion_Test_Eval*)
|
||||
This notebook introduces the first true held‑out test evaluation for an early‑fusion CNN.
|
||||
A subject‑exclusive train/test split is created before cross‑validation.
|
||||
|
||||
The intention behind this notebook is to:
|
||||
- Provide a deployment‑realistic performance estimate
|
||||
- Compare validation‑fold behavior with true test‑set behavior
|
||||
- Visualize ROC and PR curves for threshold analysis
|
||||
|
||||
| Metric / Model | CNN_crossVal_EarlyFusion_Test_Eval |
|
||||
|----------------|-------------------------------------|
|
||||
| Test Accuracy | 0.913 |
|
||||
| Test F1 | 0.927 |
|
||||
| Test AUC | 0.967 |
|
||||
| Balanced Accuracy | 0.907 |
|
||||
| Precision | 0.918 |
|
||||
| Recall | 0.937 |
|
||||
|
||||
#### Confusion Matrix
|
||||

|
||||
|
||||
*Figure 4.1.8.1: Confusion matrix of the Early‑Fusion model.*
|
||||
|
||||
#### ROC-Curve
|
||||

|
||||
|
||||
*Figure 4.1.8.2: ROC-Curve of the Early‑Fusion model.*
|
||||
|
||||
### 4.1.9 Hybrid‑Fusion CNN with Independent Test Evaluation (Notebook: *CNN_crossVal_HybridFusion_Test_Eval*)
|
||||
This notebook extends hybrid fusion with a subject‑exclusive train/test split and full test‑set evaluation.
|
||||
|
||||
The intention behind this notebook is to:
|
||||
- Evaluate hybrid fusion under realistic deployment conditions
|
||||
- Compare hybrid vs. early fusion on unseen subjects
|
||||
- Provide full diagnostic plots (ROC, PR, confusion matrices)
|
||||
|
||||
| Metric / Model | CNN_crossVal_HybridFusion_Test_Eval |
|
||||
|----------------|--------------------------------------|
|
||||
| Test Accuracy | 0.950 |
|
||||
| Test F1 | 0.959 |
|
||||
| Test AUC | 0.983 |
|
||||
| Balanced Accuracy | 0.942 |
|
||||
| Precision | 0.933 |
|
||||
| Recall | 0.986 |
|
||||
|
||||
#### Confusion Matrix
|
||||

|
||||
|
||||
*Figure 4.1.9.1: Confusion matrix of the Hybrid‑Fusion model.*
|
||||
|
||||
#### ROC-Curve
|
||||

|
||||
|
||||
*Figure 4.1.9.2: ROC-Curve of the Hybrid‑Fusion model.*
|
||||
|
||||
### 4.1.10 Summary
|
||||
Across all nine notebooks, the project progresses from a simple AU‑only baseline to advanced multimodal hybrid‑fusion architectures with independent test evaluation.
|
||||
|
||||
The final experiments revealed that hybrid fusion provides a measurable performance advantage over early fusion. While both approaches achieve strong results, the hybrid‑fusion model reaches higher overall accuracy (95% vs. 91.3%) and substantially stronger recall (98.6% vs. 93.7%), indicating that it is more effective at correctly identifying high‑workload samples.
|
||||
Early fusion, however, shows slightly better precision, suggesting that it produces fewer false positives.
|
||||
|
||||
Looking ahead, further improvements could likely be achieved through more extensive hyperparameter tuning, as the current results suggest that additional optimization headroom remains.
|
||||
|
||||
### 4.2 XGBoost
|
||||
This documentation outlines the evolution of the XGBoost classification pipeline for cognitive workload detection. The project transitioned from a basic unimodal setup to a sophisticated, multi-stage hybrid system incorporating advanced statistical filtering and deep feature extraction.
|
||||
During the model creation several methods were used to improve the model accuracy. During training, the biggest challenge was always the high overfitting of the model. Even in the last version with explicit regulation parameters the overall accuracy couldn't be improved more than the different methods before.
|
||||
The model overall was not that good, as the highest accuracy we could achieve was around 65%, which is a bit higher than Fraunhofer achieved in the ADABase-Paper.
|
||||
|
||||
### 4.2.1 Classical XGBoost Baseline
|
||||
|
||||
To establish a performance baseline, a classical Extreme Gradient Boosting (XGBoost) model was implemented. XGBoost was selected for its ability to handle non-linear relationships and its inherent regularization, which helps prevent overfitting in high-dimensional feature spaces like Facial Action Units. XGBoost was picked because of its usage in the ADABase Paper. Initially, the model utilized raw Action Unit sums with global normalization to determine the basic predictability of workload from facial muscle activity alone.
|
||||
|
||||
| Metric / Model | Classical XGBoost |
|
||||
| --- | --- |
|
||||
| Accuracy | 0.581 |
|
||||
| AUC | 0.562 |
|
||||
| F1-Score | 0.652 |
|
||||
|
||||
### 4.2.2 XGBoost with GroupKFold Validation
|
||||
|
||||
To address the challenge of inter-subject variability, the validation strategy was upgraded to `GroupKFold`. In behavioral data, samples from the same subject are highly correlated. Standard cross-validation often leads to data leakage, where the model memorizes individual facial characteristics. By ensuring that a subject's data is never shared between the training and validation sets, this iteration provides a scientifically rigorous measure of how the model generalizes to entirely unseen individuals.
|
||||
|
||||
| Metric / Model | XGBoost (GroupKFold) |
|
||||
| --- | --- |
|
||||
| Accuracy | 0.586 |
|
||||
| AUC | 0.573 |
|
||||
| F1-Score | 0.651 |
|
||||
|
||||
### 4.2.3 Hybrid XGBoost with Autoencoder
|
||||
|
||||
To improve feature quality, a hybrid approach was introduced by pre-training a deep Autoencoder. The encoder branch was used to compress 20 raw Action Units into a 5-dimensional latent space. This non-linear dimensionality reduction aims to capture muscle synergies and filter out noise that decision trees might struggle with. The XGBoost classifier was then trained on these machine-learned representations rather than raw inputs.
|
||||
|
||||
| Metric / Model | XGBoost + Autoencoder |
|
||||
| --- | --- |
|
||||
| Accuracy | 0.589 |
|
||||
| AUC | 0.575 |
|
||||
| F1-Score | 0.650 |
|
||||
|
||||
### 4.2.4 Robust XGBoost with MAD Outlier Removal
|
||||
|
||||
Recognizing that physiological and AU data often contain sensor artifacts, a robust preprocessing layer was added using Median Absolute Deviation (MAD). Unlike standard deviation, MAD is resilient to extreme outliers. By calculating a Robust Z-score and filtering signals in the training set, the model learned from a "clean" representation of cognitive states, significantly improving the stability of the gradient boosting process.
|
||||
|
||||
| Metric / Model | XGBoost + MAD |
|
||||
| --- | --- |
|
||||
| Accuracy | 0.641 |
|
||||
| AUC | 0.610 |
|
||||
| F1-Score | 0.733 |
|
||||
|
||||
### 4.2.5 Combined Dataset of Action Units and EyeTracking
|
||||
|
||||
This iteration involved training, which refined a robust pipeline on a new, expanded dataset. This dataset integrated both high-frequency facial action units and advanced eye-tracking metrics (pupillometry and fixations).
|
||||
Since the recreation of the EyeTracking data in the lab was in doubt, only Action Units were used in the first XGBoost models. Now the model also implemented the EyeTracking data as features.
|
||||
By applying performance-based subject splitting, we ensured that the training and test sets were balanced not only by label but by the subjects' underlying skill levels, resulting in the most deployable version of the AI.
|
||||
|
||||
| Metric / Model | Final Combined Model |
|
||||
| --- | --- |
|
||||
| Accuracy | 0.659 |
|
||||
| AUC | 0.621 |
|
||||
| F1-Score | 0.715 |
|
||||
|
||||
### 4.2.6 Regularized XGBoost with Complexity Control
|
||||
|
||||
Building upon the robust preprocessing of the previous steps, this iteration focuses on strict **complexity control** within the XGBoost architecture. To mitigate the 100% training accuracy observed in earlier unimodal tests—a clear indicator of overfitting—we introduced explicit **L1 (reg_alpha)** and **L2 (reg_lambda)** regularization parameters into the GridSearch space.
|
||||
|
||||
By penalizing large weights and promoting feature sparsity, the model is forced to prioritize the most globally relevant Action Units. Furthermore, the tree depth was intentionally restricted (`max_depth`: 2-4), and an **Early Stopping** callback with a 30-round patience window was implemented. This ensures that training terminates at the point of optimal generalization, capturing the essential physiological trends of cognitive load while ignoring subject-specific noise.
|
||||
|
||||
| Metric / Model | Regularized XGBoost |
|
||||
| --- | --- |
|
||||
| Accuracy | 0.665 |
|
||||
| AUC | 0.646 |
|
||||
| F1-Score | 0.727 |
|
||||
|
||||
### 4.3 Isolation Forest
|
||||
To start with unsupervised learning techniques, `IsolationForest.ipynb`was created to research how well a simple ensemble classificator performs on the created dataset.
|
||||
The notebook comes with one class grid search for hyperparameter tuning as well as a ROC curve that allows manual fine tuning.
|
||||
Overall, our experiments have shown, that this approach is not sufficient, with the following results:
|
||||
| Metric / Model | Isolation Forest |
|
||||
|----------------|---------|
|
||||
| Best Balanced Accuracy |0.57|
|
||||
| Best AUC | 0.61|
|
||||
|
||||
In detail, the classificator tends to classify the majority of samples as low load and is therefore not sufficient to be used for later deployment.
|
||||
|
||||
|
||||
### 4.4 One Class SVM with Autoencoder
|
||||
The training of an On Class SVM on the data from the dataset resulted in every sample was predicted as an anomaly.
|
||||
In the next step, an autoencoder is pretrained to learn representation of the data. Afterwards, the encoder is used for preprocessing, which leads to OCSVM training on encoder output.
|
||||
The training includes hyperparameter tuning through gridsearch cv.
|
||||
Encoder output is visualized with print statements and plots that show the encoded data for both low and high load samples.
|
||||
We see that the encoder struggles to represent the unseen high load samples differently. As a consequence, the One Class SVM also does not achieve sufficient performance.
|
||||
| Metric / Model | One Class SVM |
|
||||
|----------------|---------|
|
||||
| Best Balanced Accuracy |0.62|
|
||||
|
||||
When the notebook is run completely, both the trained encoder and svm are saved for later use given the save paths are set correctly.
|
||||
### 4.5 Deep SVDD
|
||||
Similar to the OCSVM training, an autoencoder is used to preprocess the data before the actual Deep SVDD training. Nevertheless, the usage is partialy different. The Dee SVDD uses a pretrained encoder to fine tune it by apllying a different loss function (which results from the theoretical concept behind Deep SVDD). This means that the encoder weights are still modified in the actual Deep SVDD training.
|
||||
Also, this approach includes **hybrid fusion of modalities**. Instead of putting all features into the same input layer, the neural network is divided into two branches, that process action units and eye-tracking features separately.
|
||||
Then, after two Dense layers each, the branches are fusioned by concatenation. From there, another two Dense layers process the data.
|
||||
The decoder is not exactly similar, as the split of the modalities happens are the very end.
|
||||
To compute the total loss, loss from both modalities is combined by sum. Users are able to change loss weights. Training includes 2x2 phases, both autoencoder and later Deep SVDD are first trained with larger learning rate, then fine tuned with a smaller learning rate.
|
||||
|
||||
| Metric / Model | Deep SVDD |
|
||||
|----------------|---------|
|
||||
| Best Balanced Accuracy |0.60|
|
||||
| Best AUC | 0.57|
|
||||
|
||||
### 4.6 General information on unsupervised approaches
|
||||
As described above, the approachs didn't meet the requirements in terms of prediction performance. For all models, both MinMax-Scaling as well as Standard-Scaling was done. Also, both subjectwise and globally. Unfortunately, the differences were not that large, which may explain why preprocessing wasn't mentioned above.
|
||||
|
||||
Future research should always keep in mind while subject-wise scaling might be better for training, it makes deployment on new subjects more difficult. Our solution, as implemented in `model_training/tools/scaler.py` calculates a fallback scaler (using mean of all subjects scaling parameters).
|
||||
|
||||
## 5) Real-Time Prediction and Messaging
|
||||
|
||||
Main script:
|
||||
- `predict_pipeline/predict_sample.py`
|
||||
|
||||
Pipeline:
|
||||
- Loads runtime config (`predict_pipeline/config.yaml`)
|
||||
- Pulls latest row from SQLite
|
||||
- Replaces missing values using `fallback` map from config file - if more than 50% of values need to be replaced, the sample is dropped and "valid=False"
|
||||
- Optionally applies scaler (`.pkl`/`.joblib`) - set via config file
|
||||
- Loads model (`.keras`, `.pkl`, `.joblib`) and predicts
|
||||
- Publishes JSON payload to MQTT topic
|
||||
|
||||
Expected payload form:
|
||||
```json
|
||||
{
|
||||
"valid": true, # false only if too many signals are invalid
|
||||
"_id": 123, # this is the sample ID from the database
|
||||
"prediction": 0 # 0 for low load, 1 for high load
|
||||
}
|
||||
```
|
||||
|
||||
### 5.1 Scheduled Prediction (Linux)
|
||||
|
||||
Files:
|
||||
- `predict_pipeline/predict.service`
|
||||
- `predict_pipeline/predict.timer`
|
||||
|
||||
Role:
|
||||
- Run inference repeatedly without manual execution
|
||||
- Timer/service configuration can be customized
|
||||
|
||||
More information on how to use and interact with the system service and timer can be found in [predict_service_timer_documentation.md](/predict_pipeline/predict_service_timer_documentation.md)
|
||||
|
||||
## 5.2 Runtime Configuration
|
||||
|
||||
Primary config file:
|
||||
- `predict_pipeline/config.yaml`
|
||||
|
||||
Sections:
|
||||
- `database`: SQLite location + table + sort key
|
||||
- `model`: model path
|
||||
- `scaler`: scaler usage + path
|
||||
- `mqtt`: broker and publish format
|
||||
- `sample.columns`: expected feature order
|
||||
- `fallback`: default values for NaN replacement
|
||||
|
||||
Important:
|
||||
- The repository currently uses environment-specific absolute paths in some scripts/configs to ensure functionality on Ohm-UX driving simulator.
|
||||
|
||||
|
||||
## 5.3 Data and Feature Expectations
|
||||
|
||||
Prediction expects SQLite rows containing:
|
||||
- `_Id`
|
||||
- `start_time` - this is not yet used for either predictions or messages
|
||||
- All configured model features (AUs + eye metrics)
|
||||
|
||||
Common feature groups (similar to own dataset):
|
||||
- `FACE_AUxx_mean` columns
|
||||
- Fixation counters and duration statistics
|
||||
- Saccade count/amplitude/duration statistics
|
||||
- Blink count/duration statistics
|
||||
- Pupil mean and IPA
|
||||
|
||||
## 5.4 Create database from scratch
|
||||
To (re-)create the custom database for deployment, use `fill_db.ipynb`. Enter the path to your dataset, drop unnecessary columns and insert a subset of data with tool functions from `tools/db_helpers`
|
||||
|
||||
## 6) Installation and Dependencies
|
||||
Due to unsolvable dependency conflicts, several environemnts need to be used in the same time.
|
||||
### 6.1 Environemnt for camera handling
|
||||
The setup of a virtual environment for the camera handling is difficult due to vary dependency conflicts.
|
||||
Therefore it is necessary to create the virtual environment with every package in the specific version and each package in the specific order.
|
||||
Furthermore the environment needs to be based on Python 3.10. The specific versions and order of the packages are described int the file:
|
||||
`requirements.txt`
|
||||
|
||||
|
||||
### 6.2 Environment for predictions
|
||||
If you want to use the existing deployment on Ohm-UX driving simulator's jetson board, activate conda environment `p310_FS_TF`, a python 3.10 environment including tensorflow and all other packages required to run `predict_sample.py`
|
||||
Otherwise, as described in `readme.md: Setup`, you can use `prediction_env.yaml`, to create a new environment that fulfills the requirements.
|
||||
|
||||
## 7) Repository File Inventory
|
||||
|
||||
### Root
|
||||
|
||||
- `.gitignore` - Git ignore rules
|
||||
- `readme.md` - minimal quickstart documentation
|
||||
- `project_report.md` - full technical documentation (this file)
|
||||
- `requirements.txt` - Python dependencies
|
||||
|
||||
### Dataset Creation
|
||||
|
||||
- `dataset_creation/parquet_file_creation.py` - local files to parquet conversion
|
||||
- `dataset_creation/create_parquet_files_from_owncloud.py` - ownCloud download + parquet conversion
|
||||
- `dataset_creation/combined_feature_creation.py` - sliding-window multimodal feature generation
|
||||
- `dataset_creation/maxDist.py` - helper/statistical utility script for eye-tracking feature creation
|
||||
|
||||
#### AU Creation
|
||||
- `dataset_creation/AU_creation/pyfeat_docu.ipynb` - py-feat exploratory notes
|
||||
|
||||
#### Camera Handling
|
||||
- `dataset_creation/camera_handling/camera_stream_AU_and_ET_new.py` - current camera + AU + eye online pipeline
|
||||
- `dataset_creation/camera_handling/eyeFeature_new.py` - eye-feature extraction from gaze parquet
|
||||
- `dataset_creation/camera_handling/db_helper.py` - SQLite helper functions (camera pipeline)
|
||||
- `dataset_creation/camera_handling/camera_stream.py` - baseline camera streaming script
|
||||
- `dataset_creation/camera_handling/eyeFeature_kalibrierung.py` - eye-feature extraction with calibration
|
||||
- `dataset_creation/camera_handling/db_test.py` - DB test utility
|
||||
|
||||
### EDA
|
||||
|
||||
- `EDA/EDA.ipynb` - main EDA notebook
|
||||
- `EDA/distribution_plots.ipynb` - distribution visualization
|
||||
- `EDA/histogramms.ipynb` - histogram analysis
|
||||
- `EDA/researchOnSubjectPerformance.ipynb` - subject-level analysis
|
||||
- `EDA/owncloud_file_access.ipynb` - ownCloud exploration/access notebook
|
||||
- `EDA/calculate_replacement_values.ipynb` - fallback/median computation notebook
|
||||
- `EDA/login.yaml` - local auth/config artifact for EDA workflows
|
||||
|
||||
### Model Training
|
||||
|
||||
#### CNN
|
||||
- `model_training/CNN/CNN_simple.ipynb`
|
||||
- `model_training/CNN/CNN_crossVal.ipynb`
|
||||
- `model_training/CNN/CNN_crossVal_EarlyFusion.ipynb`
|
||||
- `model_training/CNN/CNN_crossVal_EarlyFusion_Filter.ipynb`
|
||||
- `model_training/CNN/CNN_crossVal_EarlyFusion_Test_Eval.ipynb`
|
||||
- `model_training/CNN/CNN_crossVal_faceAUs.ipynb`
|
||||
- `model_training/CNN/CNN_crossVal_faceAUs_eyeFeatures.ipynb`
|
||||
- `model_training/CNN/CNN_crossVal_HybridFusion.ipynb`
|
||||
- `model_training/CNN/CNN_crossVal_HybridFusion_Test_Eval.ipynb`
|
||||
- `model_training/CNN/deployment_pipeline.ipynb`
|
||||
|
||||
#### XGBoost
|
||||
- `model_training/xgboost/xgboost.ipynb`
|
||||
- `model_training/xgboost/xgboost_groupfold.ipynb`
|
||||
- `model_training/xgboost/xgboost_new_dataset.ipynb`
|
||||
- `model_training/xgboost/xgboost_regulated.ipynb`
|
||||
- `model_training/xgboost/xgboost_with_AE.ipynb`
|
||||
- `model_training/xgboost/xgboost_with_MAD.ipynb`
|
||||
|
||||
#### Isolation Forest
|
||||
- `model_training/IsolationForest/iforest_training.ipynb`
|
||||
|
||||
#### OCSVM
|
||||
- `model_training/OCSVM/ocsvm_with_AE.ipynb`
|
||||
|
||||
#### DeepSVDD
|
||||
- `model_training/DeepSVDD/deepSVDD.ipynb`
|
||||
|
||||
#### MAD Outlier Removal
|
||||
- `model_training/MAD_outlier_removal/mad_outlier_removal.ipynb`
|
||||
- `model_training/MAD_outlier_removal/mad_outlier_removal_median.ipynb`
|
||||
|
||||
#### Shared Training Tools
|
||||
- `model_training/tools/scaler.py`
|
||||
- `model_training/tools/performance_split.py`
|
||||
- `model_training/tools/mad_outlier_removal.py`
|
||||
- `model_training/tools/evaluation_tools.py`
|
||||
|
||||
### Prediction Pipeline
|
||||
|
||||
- `predict_pipeline/predict_sample.py` - runtime prediction + MQTT publish
|
||||
- `predict_pipeline/config.yaml` - runtime database/model/scaler/mqtt config
|
||||
- `predict_pipeline/fill_db.ipynb` - helper notebook for DB setup/testing
|
||||
- `predict_pipeline/predict.service` - systemd service unit
|
||||
- `predict_pipeline/predict.timer` - systemd timer unit
|
||||
- `predict_pipeline/predict_service_timer_documentation.md` - Linux service/timer guide
|
||||
|
||||
### Generic Tools
|
||||
|
||||
- `tools/db_helpers.py` - common SQLite utilities used to get newest sample for prediction
|
||||
|
||||
## 8) Additional Information
|
||||
|
||||
- Several paths are hardcoded on purpose to ensure compability with the jetsonboard at the OHM-UX driving simulator.
|
||||
- Camera and AU processing are resource-intensive; version pinning and hardware validation are recommended.
|
||||
- To access our dataset and other valuables, the `data-paulusjafahrsimulator-gpu`
|
||||
directory contains the raw data, performance results, parquet files, and the final dataset.
|
||||
@@ -1,67 +1,54 @@
|
||||
# Multimodal Driver State Analysis
|
||||
|
||||
Ein umfassendes Framework zur Analyse von Fahrerverhalten durch kombinierte Feature-Extraktion aus Facial Action Units (AU) und Eye-Tracking Daten.
|
||||
Short overview: this repository contains the data, feature, training, and inference pipeline for multimodal driver-state analysis using facial AUs and eye-tracking signals.
|
||||
|
||||
## 📋 Projektübersicht
|
||||
For full documentation, see [project_report.md](project_report.md).
|
||||
|
||||
Dieses Projekt verarbeitet multimodale Sensordaten aus Fahrsimulator-Studien und extrahiert zeitbasierte Features für die Analyse von Fahrerzuständen. Die Pipeline kombiniert:
|
||||
## Quickstart
|
||||
|
||||
- **Facial Action Units (AU)**: 20 Gesichtsaktionseinheiten zur Emotionserkennung
|
||||
- **Eye-Tracking**: Fixationen, Sakkaden, Blinks und Pupillenmetriken
|
||||
|
||||
## 🎯 Features
|
||||
|
||||
### Datenverarbeitung
|
||||
- **Sliding Window Aggregation**: 50-Sekunden-Fenster mit 5-Sekunden-Schrittweite
|
||||
- **Hierarchische Gruppierung**: Automatische Segmentierung nach STUDY/LEVEL/PHASE
|
||||
- **Robuste Fehlerbehandlung**: Graceful Degradation bei fehlenden Modalitäten
|
||||
|
||||
### Extrahierte Features
|
||||
|
||||
#### Facial Action Units (20 AUs)
|
||||
Für jede AU wird der Mittelwert pro Window berechnet:
|
||||
- AU01 (Inner Brow Raiser) bis AU43 (Eyes Closed)
|
||||
- Aggregation: `mean` über 50s Window
|
||||
|
||||
#### Eye-Tracking Features
|
||||
**Fixationen:**
|
||||
- Anzahl nach Dauer-Kategorien (66-150ms, 300-500ms, >1000ms, >100ms)
|
||||
- Mittelwert und Median der Fixationsdauer
|
||||
|
||||
**Sakkaden:**
|
||||
- Anzahl, mittlere Amplitude, mittlere/mediane Dauer
|
||||
|
||||
**Blinks:**
|
||||
- Anzahl, mittlere/mediane Dauer
|
||||
|
||||
**Pupille:**
|
||||
- Mittlere Pupillengröße
|
||||
- Index of Pupillary Activity (IPA) - Hochfrequenzkomponente (0.6-2.0 Hz)
|
||||
|
||||
## 🏗️ Projektstruktur
|
||||
|
||||
to be continued.
|
||||
|
||||
## 🚀 Installation
|
||||
|
||||
### Voraussetzungen
|
||||
### 1) Setup
|
||||
Activate the conda-repository "camera_stream_AU_ET_test".
|
||||
```bash
|
||||
Python 3.12
|
||||
conda activate camera_stream_AU_ET_test
|
||||
```
|
||||
**Make sure, another environment that fulfills prediction_env.yaml is available**, matching with predict_pipeline/predict.service
|
||||
See `predict_pipeline/predict_service_timer_documentation.md`
|
||||
to get an overview over all available conda environments on your device, use this command in anaconda prompt terminal:
|
||||
```bash
|
||||
conda info --envs
|
||||
```
|
||||
Optionally, create a new environment based on the yaml-file:
|
||||
```bash
|
||||
conda env create -f prediction_env.yaml
|
||||
```
|
||||
Ohm-UX driving simulator jetson board only: The conda-environment `p310_FS_TF` is used for predictions.
|
||||
|
||||
### 2) Camera AU + Eye Pipeline (`camera_stream_AU_and_ET_new.py`)
|
||||
|
||||
1. Open `dataset_creation/camera_handling/camera_stream_AU_and_ET_new.py` and adjust:
|
||||
- `DB_PATH`
|
||||
- `CAMERA_INDEX`
|
||||
- `OUTPUT_DIR` (optional)
|
||||
|
||||
2. Start camera capture and feature extraction:
|
||||
```bash
|
||||
python dataset_creation/camera_handling/camera_stream_AU_and_ET_new.py
|
||||
```
|
||||
|
||||
### Dependencies
|
||||
3. Stop with `q` in the camera window.
|
||||
|
||||
|
||||
### 3) Predict Pipeline (`predict_pipeline/predict_sample.py`)
|
||||
|
||||
1. Edit `predict_pipeline/config.yaml` and set:
|
||||
- `database.path`, `database.table`, `database.key`
|
||||
- `model.path`
|
||||
- `scaler.path` (if `use_scaling: true`)
|
||||
- MQTT settings under `mqtt`
|
||||
|
||||
2. Run one prediction cycle:
|
||||
```bash
|
||||
pip install -r requirements.txt
|
||||
python predict_pipeline/predict_sample.py
|
||||
```
|
||||
|
||||
**Wichtigste Pakete:**
|
||||
- `pandas`, `numpy` - Datenverarbeitung
|
||||
- `scipy` - Signalverarbeitung
|
||||
- `scikit-learn` - Feature-Skalierung & ML
|
||||
- `pygazeanalyser` - Eye-Tracking Analyse
|
||||
- `pyarrow` - Parquet I/O
|
||||
|
||||
## 💻 Usage
|
||||
|
||||
### 1. Feature-Extraktion
|
||||
to be continued
|
||||
3. Use [predict_service_timer_documentation.md](/predict_pipeline/predict_service_timer_documentation.md) to see how to use the service and timer for automation. On Ohm-UX driving simulator's jetson board, the service runs in the background and starts automatically when the device is booting.
|
||||
@@ -0,0 +1,27 @@
|
||||
# Core data + ML utilities
|
||||
numpy
|
||||
pandas
|
||||
scipy
|
||||
scikit-learn
|
||||
pyarrow
|
||||
joblib
|
||||
PyYAML
|
||||
matplotlib
|
||||
|
||||
# Prediction pipeline
|
||||
paho-mqtt
|
||||
tensorflow
|
||||
|
||||
# Camera / feature extraction stack
|
||||
# It is necessary to create an extra environment for the extraction pipeline, because of the different version needs
|
||||
Python==3.10
|
||||
numpy==1.24.4
|
||||
scipy==1.10.1
|
||||
opencv-python==4.7.0.72
|
||||
mediapipe==0.10.13
|
||||
py-feat
|
||||
|
||||
# Data ingestion (ownCloud + HDF)
|
||||
pyocclient
|
||||
h5py
|
||||
tables
|
||||
@@ -0,0 +1,166 @@
|
||||
import os
|
||||
import sqlite3
|
||||
import pandas as pd
|
||||
|
||||
|
||||
def connect_db(path_to_file: os.PathLike) -> tuple[sqlite3.Connection, sqlite3.Cursor]:
|
||||
''' Establishes a connection with a sqlite3 database. '''
|
||||
conn = sqlite3.connect(path_to_file)
|
||||
cursor = conn.cursor()
|
||||
return conn, cursor
|
||||
|
||||
def disconnect_db(conn: sqlite3.Connection, cursor: sqlite3.Cursor, commit: bool = True) -> None:
|
||||
''' Commits all remaining changes and closes the connection with an sqlite3 database. '''
|
||||
cursor.close()
|
||||
if commit: conn.commit() # commit all pending changes made to the sqlite3 database before closing
|
||||
conn.close()
|
||||
|
||||
def create_table(
|
||||
conn: sqlite3.Connection,
|
||||
cursor: sqlite3.Cursor,
|
||||
table_name: str,
|
||||
columns: dict,
|
||||
constraints: dict,
|
||||
primary_key: dict,
|
||||
commit: bool = True
|
||||
) -> str:
|
||||
'''
|
||||
Creates a new empty table with the given columns, constraints and primary key.
|
||||
|
||||
:param columns: dict with column names (=keys) and dtypes (=values) (e.g. BIGINT, INT, ...)
|
||||
:param constraints: dict with column names (=keys) and list of constraints (=values) (like [\'NOT NULL\'(,...)])
|
||||
:param primary_key: dict with primary key name (=key) and list of attributes which combined define the table's primary key (=values, like [\'att1\'(,...)])
|
||||
'''
|
||||
assert len(primary_key.keys()) == 1
|
||||
sql = f'CREATE TABLE {table_name} (\n '
|
||||
for column,dtype in columns.items():
|
||||
sql += f'{column} {dtype}{" "+" ".join(constraints[column]) if column in constraints.keys() else ""},\n '
|
||||
if list(primary_key.keys())[0]: sql += f'CONSTRAINT {list(primary_key.keys())[0]} '
|
||||
sql += f'PRIMARY KEY ({", ".join(list(primary_key.values())[0])})\n)'
|
||||
cursor.execute(sql)
|
||||
if commit: conn.commit()
|
||||
return sql
|
||||
|
||||
def add_columns_to_table(
|
||||
conn: sqlite3.Connection,
|
||||
cursor: sqlite3.Cursor,
|
||||
table_name: str,
|
||||
columns: dict,
|
||||
constraints: dict = dict(),
|
||||
commit: bool = True
|
||||
) -> str:
|
||||
''' Adds one/multiple columns (each with a list of constraints) to the given table. '''
|
||||
sql_total = ''
|
||||
for column,dtype in columns.items(): # sqlite can only add one column per query
|
||||
sql = f'ALTER TABLE {table_name}\n '
|
||||
sql += f'ADD "{column}" {dtype}{" "+" ".join(constraints[column]) if column in constraints.keys() else ""}'
|
||||
sql_total += sql + '\n'
|
||||
cursor.execute(sql)
|
||||
if commit: conn.commit()
|
||||
return sql_total
|
||||
|
||||
|
||||
|
||||
|
||||
def insert_rows_into_table(
|
||||
conn: sqlite3.Connection,
|
||||
cursor: sqlite3.Cursor,
|
||||
table_name: str,
|
||||
columns: dict,
|
||||
commit: bool = True
|
||||
) -> str:
|
||||
'''
|
||||
Inserts values as multiple rows into the given table.
|
||||
|
||||
:param columns: dict with column names (=keys) and values to insert as lists with at least one element (=values)
|
||||
|
||||
Note: The number of given values per attribute must match the number of rows to insert!
|
||||
Note: The values for the rows must be of normal python types (e.g. list, str, int, ...) instead of e.g. numpy arrays!
|
||||
'''
|
||||
assert len(set(map(len, columns.values()))) == 1, 'ERROR: Provide equal number of values for each column!'
|
||||
assert len(set(list(map(type,columns.values())))) == 1 and isinstance(list(columns.values())[0], list), 'ERROR: Provide values as Python lists!'
|
||||
assert set([type(a) for b in list(columns.values()) for a in b]).issubset({str,int,float,bool}), 'ERROR: Provide values as basic Python data types!'
|
||||
|
||||
values = list(zip(*columns.values()))
|
||||
sql = f'INSERT INTO {table_name} ({", ".join(columns.keys())})\n VALUES ({("?,"*len(values[0]))[:-1]})'
|
||||
cursor.executemany(sql, values)
|
||||
if commit: conn.commit()
|
||||
return sql
|
||||
|
||||
def update_multiple_rows_in_table(
|
||||
conn: sqlite3.Connection,
|
||||
cursor: sqlite3.Cursor,
|
||||
table_name: str,
|
||||
new_vals: dict,
|
||||
conditions: str,
|
||||
commit: bool = True
|
||||
) -> str:
|
||||
'''
|
||||
Updates attribute values of some rows in the given table.
|
||||
|
||||
:param new_vals: dict with column names (=keys) and the new values to set (=values)
|
||||
:param conditions: string which defines all concatenated conditions (e.g. \'cond1 AND (cond2 OR cond3)\' with cond1: att1=5, ...)
|
||||
'''
|
||||
assignments = ', '.join([f'{k}={v}' for k,v in zip(new_vals.keys(), new_vals.values())])
|
||||
sql = f'UPDATE {table_name}\n SET {assignments}\n WHERE {conditions}'
|
||||
cursor.execute(sql)
|
||||
if commit: conn.commit()
|
||||
return sql
|
||||
|
||||
def delete_rows_from_table(
|
||||
conn: sqlite3.Connection,
|
||||
cursor: sqlite3.Cursor,
|
||||
table_name: str,
|
||||
conditions: str,
|
||||
commit: bool = True
|
||||
) -> str:
|
||||
'''
|
||||
Deletes rows from the given table.
|
||||
|
||||
:param conditions: string which defines all concatenated conditions (e.g. \'cond1 AND (cond2 OR cond3)\' with cond1: att1=5, ...)
|
||||
'''
|
||||
sql = f'DELETE FROM {table_name} WHERE {conditions}'
|
||||
cursor.execute(sql)
|
||||
if commit: conn.commit()
|
||||
return sql
|
||||
|
||||
|
||||
|
||||
def get_data_from_table(
|
||||
conn: sqlite3.Connection,
|
||||
table_name: str,
|
||||
columns_list: list = ['*'],
|
||||
aggregations: [None,dict] = None,
|
||||
where_conditions: [None,str] = None,
|
||||
order_by: [None, dict] = None,
|
||||
limit: [None, int] = None,
|
||||
offset: [None, int] = None
|
||||
) -> pd.DataFrame:
|
||||
'''
|
||||
Helper function which returns (if desired: aggregated) contents from the given table as a pandas DataFrame. The rows can be filtered by providing the condition as a string.
|
||||
|
||||
:param columns_list: use if no aggregation is needed to select which columns to get from the table
|
||||
:param (optional) aggregations: use to apply aggregations on the data from the table; dictionary with column(s) as key(s) and aggregation(s) as corresponding value(s) (e.g. {'col1': 'MIN', 'col2': 'AVG', ...} or {'*': 'COUNT'})
|
||||
:param (optional) where_conditions: string which defines all concatenated conditions (e.g. \'cond1 AND (cond2 OR cond3)\' with cond1: att1=5, ...) applied on table.
|
||||
:param (optional) order_by: dict defining the ordering of the outputs with column(s) as key(s) and ordering as corresponding value(s) (e.g. {'col1': 'ASC'})
|
||||
:param (optional) limit: use to limit the number of returned rows
|
||||
:param (optional) offset: use to skip the first n rows before displaying
|
||||
|
||||
Note: If aggregations is set, the columns_list is ignored.
|
||||
Note: Get all data as a DataFrame with get_data_from_table(conn, table_name).
|
||||
Note: If one output is wanted (e.g. count(*) or similar), get it with get_data_from_table(...).iloc[0,0] from the DataFrame.
|
||||
'''
|
||||
assert columns_list or aggregations
|
||||
|
||||
if aggregations:
|
||||
selection = [f'{agg}({col})' for col,agg in aggregations.items()]
|
||||
else:
|
||||
selection = columns_list
|
||||
selection = ", ".join(selection)
|
||||
where_conditions = 'WHERE ' + where_conditions if where_conditions else ''
|
||||
order_by = 'ORDER BY ' + ', '.join([f'{k} {v}' for k,v in order_by.items()]) if order_by else ''
|
||||
limit = f'LIMIT {limit}' if limit else ''
|
||||
offset = f'OFFSET {offset}' if offset else ''
|
||||
|
||||
sql = f'SELECT {selection} FROM {table_name} {where_conditions} {order_by} {limit} {offset}'
|
||||
return pd.read_sql_query(sql, conn)
|
||||
Reference in New Issue
Block a user