{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":41880,"databundleVersionId":5677426,"sourceType":"competition"},{"sourceId":105488,"sourceType":"modelInstanceVersion","modelInstanceId":88399,"modelId":112626},{"sourceId":105489,"sourceType":"modelInstanceVersion","modelInstanceId":88400,"modelId":112627}],"dockerImageVersionId":30762,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Основная задача на ближайшие недели - попробовать сравнить разные подходы (XGBoost, LightGBM, CNN, LSTM) на одних и тех же обучающих, валидационных и тестовых выборках.\nКаждый метод отдельно удалось изучить, теперь нужно понять, какой метод дает лучшие метрики при условии обучения и тестирования на одних и тех же данных.\nНапример, сначала разбить данные (выбрав наиболее удобный для этого датасет, например, tDCS FOG), сохранить это разбиение. А затем работать с ним в каждом блокноте, обучить модели на одной и той же обучающей выборке и рассчитать одни и те же метрики на тестовой выборке.","metadata":{}},{"cell_type":"markdown","source":"接下来几周的主要任务是尝试在相同的训练、验证和测试数据集上比较不同的方法（XGBoost、LightGBM、CNN、LSTM）。\n每种方法都已经单独研究过了，现在需要了解在相同数据上进行训练和测试时，哪种方法能给出更好的指标。\n例如，首先拆分数据（选择一个最适合的数据集，例如 tDCS FOG），保存这个拆分。然后在每个笔记本中使用它，在相同的训练集上训练模型，并在测试集上计算相同的指标。","metadata":{}},{"cell_type":"markdown","source":"'#'+data 表示更新时间","metadata":{}},{"cell_type":"markdown","source":"================== 1. Загрузка данных и оптимизация памяти ==================","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport gc\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder, StandardScaler\nfrom sklearn.utils.class_weight import compute_class_weight\nimport joblib  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T19:36:14.138629Z","iopub.execute_input":"2025-02-25T19:36:14.138877Z","iopub.status.idle":"2025-02-25T19:36:19.416334Z","shell.execute_reply.started":"2025-02-25T19:36:14.13885Z","shell.execute_reply":"2025-02-25T19:36:19.415333Z"}},"outputs":[],"execution_count":1},{"cell_type":"code","source":"def reduce_memory_usage(df):\n    \"\"\"Оптимизация памяти\"\"\"\n    start_mem = df.memory_usage().sum() / 1024**2\n    for col in df.columns:\n        col_type = df[col].dtype.name\n        if col_type not in ['datetime64[ns]', 'category']:\n            if col_type != 'object':\n                c_min, c_max = df[col].min(), df[col].max()\n                if 'int' in str(col_type):\n                    int_types = [np.int8, np.int16, np.int32, np.int64]\n                    for it in int_types:\n                        if c_min > np.iinfo(it).min and c_max < np.iinfo(it).max:\n                            df[col] = df[col].astype(it)\n                            break\n                else:\n                    float_types = [np.float16, np.float32]\n                    for ft in float_types:\n                        if c_min > np.finfo(ft).min and c_max < np.finfo(ft).max:\n                            df[col] = df[col].astype(ft)\n                            break\n            else:\n                df[col] = df[col].astype('category')\n    mem_usg = df.memory_usage().sum() / 1024**2 \n    print(f\"内存优化: {start_mem:.2f} MB -> {mem_usg:.2f} MB\")\n    return df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T19:37:27.777668Z","iopub.execute_input":"2025-02-25T19:37:27.77859Z","iopub.status.idle":"2025-02-25T19:37:27.788886Z","shell.execute_reply.started":"2025-02-25T19:37:27.77854Z","shell.execute_reply":"2025-02-25T19:37:27.788075Z"}},"outputs":[],"execution_count":2},{"cell_type":"code","source":"# Загрузка и объединение данных\nDATA_ROOT_TDCSFOG = '/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/tdcsfog/'\ntdcsfog = pd.concat([pd.read_csv(os.path.join(root, name)).assign(file=name.split('.')[0]) \n                    for root, _, files in os.walk(DATA_ROOT_TDCSFOG) \n                    for name in files], axis=0)\ntdcsfog = reduce_memory_usage(tdcsfog)\n\ntdcsfog_metadata = pd.read_csv(\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/tdcsfog_metadata.csv\")\ntdcsfog_m = tdcsfog_metadata.merge(tdcsfog, how='inner', left_on='Id', right_on='file').drop('file', axis=1)\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T19:40:03.242251Z","iopub.execute_input":"2025-02-25T19:40:03.242581Z","iopub.status.idle":"2025-02-25T19:40:14.168676Z","shell.execute_reply.started":"2025-02-25T19:40:03.242555Z","shell.execute_reply":"2025-02-25T19:40:14.167799Z"}},"outputs":[{"name":"stdout","text":"内存优化: 484.96 MB -> 154.95 MB\n","output_type":"stream"},{"execution_count":4,"output_type":"execute_result","data":{"text/plain":"0"},"metadata":{}}],"execution_count":4},{"cell_type":"markdown","source":"================== 2. Инженерия меток ==================","metadata":{}},{"cell_type":"code","source":"# Проверка конфликтов меток\nconflict_mask = (tdcsfog_m[['StartHesitation', 'Turn', 'Walking']].sum(axis=1) > 1)\ntdcsfog_m = tdcsfog_m[~conflict_mask]\n\n# Генерация многоклассовых меток\nconditions = [\n    (tdcsfog_m['StartHesitation'] == 1),\n    (tdcsfog_m['Turn'] == 1),\n    (tdcsfog_m['Walking'] == 1)\n]\ntdcsfog_m['event'] = np.select(conditions, ['StartHesitation', 'Turn', 'Walking'], default='Normal')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:00:50.931194Z","iopub.execute_input":"2025-02-25T20:00:50.931539Z","iopub.status.idle":"2025-02-25T20:00:53.499446Z","shell.execute_reply.started":"2025-02-25T20:00:50.931512Z","shell.execute_reply":"2025-02-25T20:00:53.49858Z"}},"outputs":[],"execution_count":8},{"cell_type":"markdown","source":"================== 3. Стратифицированное разделение набора данных ==================","metadata":{}},{"cell_type":"code","source":"features = ['AccV', 'AccML', 'AccAP']\nle = LabelEncoder()\ntdcsfog_m['target'] = le.fit_transform(tdcsfog_m['event'])\n\n# Стратифицированное разделение набора данных (60% обучение, 20% валидация, 20% тестирование)\nX_temp, X_test, y_temp, y_test = train_test_split(\n    tdcsfog_m[features], tdcsfog_m['target'],\n    test_size=0.2, stratify=tdcsfog_m['target'], random_state=1004\n)\nX_train, X_val, y_train, y_val = train_test_split(\n    X_temp, y_temp, test_size=0.25, stratify=y_temp, random_state=1004\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:00:53.735635Z","iopub.execute_input":"2025-02-25T20:00:53.736341Z","iopub.status.idle":"2025-02-25T20:01:00.345658Z","shell.execute_reply.started":"2025-02-25T20:00:53.736311Z","shell.execute_reply":"2025-02-25T20:01:00.344951Z"}},"outputs":[],"execution_count":9},{"cell_type":"markdown","source":"================== 4. Расчет весов выборки ==================","metadata":{}},{"cell_type":"code","source":"class_weights = compute_class_weight('balanced', classes=np.unique(y_train), y=y_train)\nsample_weights = np.array([class_weights[y] for y in y_train])  # 为每个样本分配权重","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:04:12.268398Z","iopub.execute_input":"2025-02-25T20:04:12.269274Z","iopub.status.idle":"2025-02-25T20:04:14.090835Z","shell.execute_reply.started":"2025-02-25T20:04:12.269238Z","shell.execute_reply":"2025-02-25T20:04:14.0898Z"}},"outputs":[],"execution_count":25},{"cell_type":"markdown","source":"================== 5. Стандартизация обработки ==================","metadata":{}},{"cell_type":"code","source":"scaler = StandardScaler()\nX_train = scaler.fit_transform(X_train)  # 覆盖原始X_train\nX_val = scaler.transform(X_val)\nX_test = scaler.transform(X_test)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:01:13.504333Z","iopub.execute_input":"2025-02-25T20:01:13.505051Z","iopub.status.idle":"2025-02-25T20:01:14.272886Z","shell.execute_reply.started":"2025-02-25T20:01:13.505015Z","shell.execute_reply":"2025-02-25T20:01:14.272227Z"}},"outputs":[],"execution_count":11},{"cell_type":"code","source":"#На самом деле это одно и то же. \nscaler = StandardScaler()\nX_train_scaled = scaler.fit_transform(X_train)\nX_val_scaled = scaler.transform(X_val)\nX_test_scaled = scaler.transform(X_test)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-27T10:29:13.462609Z","iopub.execute_input":"2025-02-27T10:29:13.462927Z","iopub.status.idle":"2025-02-27T10:29:13.761218Z","shell.execute_reply.started":"2025-02-27T10:29:13.462887Z","shell.execute_reply":"2025-02-27T10:29:13.759774Z"},"_kg_hide-output":false},"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mNameError\u001b[0m                                 Traceback (most recent call last)","Cell \u001b[0;32mIn[1], line 1\u001b[0m\n\u001b[0;32m----> 1\u001b[0m scaler \u001b[38;5;241m=\u001b[39m \u001b[43mStandardScaler\u001b[49m()\n\u001b[1;32m      2\u001b[0m X_train_scaled \u001b[38;5;241m=\u001b[39m scaler\u001b[38;5;241m.\u001b[39mfit_transform(X_train)\n\u001b[1;32m      3\u001b[0m X_val_scaled \u001b[38;5;241m=\u001b[39m scaler\u001b[38;5;241m.\u001b[39mtransform(X_val)\n","\u001b[0;31mNameError\u001b[0m: name 'StandardScaler' is not defined"],"ename":"NameError","evalue":"name 'StandardScaler' is not defined","output_type":"error"}],"execution_count":1},{"cell_type":"markdown","source":"================= 6. Сохранение набора данных ==================","metadata":{}},{"cell_type":"code","source":"# 创建保存目录 \n!mkdir -p /kaggle/working/datasets\n!mkdir -p /kaggle/working/preprocessors","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:01:16.151644Z","iopub.execute_input":"2025-02-25T20:01:16.152464Z","iopub.status.idle":"2025-02-25T20:01:18.189878Z","shell.execute_reply.started":"2025-02-25T20:01:16.152433Z","shell.execute_reply":"2025-02-25T20:01:18.188772Z"}},"outputs":[],"execution_count":12},{"cell_type":"code","source":"#保存特征数据（带列名）\ndef save_dataset(df, path):\n    \"\"\"优化内存保存为Parquet格式\"\"\"\n    df = df.copy()\n    # 优化内存（可选）\n    for col in df.columns:\n        if df[col].dtype == 'float64':\n            df[col] = df[col].astype('float32')\n    # 保存\n    df.to_parquet(path, index=False)\n    print(f\"Saved {path} ({df.shape})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:01:19.118321Z","iopub.execute_input":"2025-02-25T20:01:19.119108Z","iopub.status.idle":"2025-02-25T20:01:19.124409Z","shell.execute_reply.started":"2025-02-25T20:01:19.119073Z","shell.execute_reply":"2025-02-25T20:01:19.123576Z"}},"outputs":[],"execution_count":13},{"cell_type":"code","source":"# Обучающий набор\nX_train_df = pd.DataFrame(X_train, columns=['AccV', 'AccML', 'AccAP'])\nsave_dataset(X_train_df, '/kaggle/working/datasets/X_train.parquet')\n\n# Валидационный набор\nX_val_df = pd.DataFrame(X_val, columns=['AccV', 'AccML', 'AccAP'])\nsave_dataset(X_val_df, '/kaggle/working/datasets/X_val.parquet')\n\n# Тестовый набор\nX_test_df = pd.DataFrame(X_test, columns=['AccV', 'AccML', 'AccAP'])\nsave_dataset(X_test_df, '/kaggle/working/datasets/X_test.parquet')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:46:30.23823Z","iopub.execute_input":"2025-02-25T20:46:30.239108Z","iopub.status.idle":"2025-02-25T20:46:31.461857Z","shell.execute_reply.started":"2025-02-25T20:46:30.239073Z","shell.execute_reply":"2025-02-25T20:46:31.460801Z"}},"outputs":[{"name":"stdout","text":"Saved /kaggle/working/datasets/X_train.parquet ((4237602, 3))\nSaved /kaggle/working/datasets/X_val.parquet ((1412535, 3))\nSaved /kaggle/working/datasets/X_test.parquet ((1412535, 3))\n","output_type":"stream"}],"execution_count":48},{"cell_type":"code","source":"#Сохранение данных меток\ndef save_labels(y, path):\n    pd.Series(y).to_csv(path, index=False, header=['label'])\n    print(f\"Saved {path} ({len(y)} samples)\")\n\nsave_labels(y_train, '/kaggle/working/datasets/y_train.csv')\nsave_labels(y_val, '/kaggle/working/datasets/y_val.csv')\nsave_labels(y_test, '/kaggle/working/datasets/y_test.csv')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:01:25.533377Z","iopub.execute_input":"2025-02-25T20:01:25.533677Z","iopub.status.idle":"2025-02-25T20:01:27.483507Z","shell.execute_reply.started":"2025-02-25T20:01:25.533653Z","shell.execute_reply":"2025-02-25T20:01:27.482507Z"}},"outputs":[{"name":"stdout","text":"Saved /kaggle/working/datasets/y_train.csv (4237602 samples)\nSaved /kaggle/working/datasets/y_val.csv (1412535 samples)\nSaved /kaggle/working/datasets/y_test.csv (1412535 samples)\n","output_type":"stream"}],"execution_count":15},{"cell_type":"code","source":"#Сохранение файла информации о наборе данных\ndataset_info = {\n    'feature_columns': ['AccV', 'AccML', 'AccAP'],\n    'label_mapping': dict(zip(le.classes_, range(len(le.classes_)))),\n    'num_classes': len(le.classes_),\n    'class_weights': class_weights.tolist() \n}\n\njoblib.dump(dataset_info, '/kaggle/working/datasets/dataset_info.joblib')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:09:51.424661Z","iopub.execute_input":"2025-02-25T20:09:51.425076Z","iopub.status.idle":"2025-02-25T20:09:51.432553Z","shell.execute_reply.started":"2025-02-25T20:09:51.425042Z","shell.execute_reply":"2025-02-25T20:09:51.431623Z"}},"outputs":[{"execution_count":29,"output_type":"execute_result","data":{"text/plain":"['/kaggle/working/datasets/dataset_info.joblib']"},"metadata":{}}],"execution_count":29},{"cell_type":"code","source":"!mkdir -p /kaggle/working/preprocessors\njoblib.dump(scaler, '/kaggle/working/preprocessors/standard_scaler.joblib')\njoblib.dump(le, '/kaggle/working/preprocessors/label_encoder.joblib')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:09:57.241687Z","iopub.execute_input":"2025-02-25T20:09:57.242563Z","iopub.status.idle":"2025-02-25T20:09:58.279951Z","shell.execute_reply.started":"2025-02-25T20:09:57.242526Z","shell.execute_reply":"2025-02-25T20:09:58.278922Z"}},"outputs":[{"name":"stderr","text":"/opt/conda/lib/python3.10/pty.py:89: RuntimeWarning: os.fork() was called. os.fork() is incompatible with multithreaded code, and JAX is multithreaded, so this will likely lead to a deadlock.\n  pid, fd = os.forkpty()\n","output_type":"stream"},{"execution_count":30,"output_type":"execute_result","data":{"text/plain":"['/kaggle/working/preprocessors/label_encoder.joblib']"},"metadata":{}}],"execution_count":30},{"cell_type":"code","source":"#Проверка результатов сохранения\n!ls -lh /kaggle/working/datasets\n!ls -lh /kaggle/working/preprocessors","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:10:05.358545Z","iopub.execute_input":"2025-02-25T20:10:05.358973Z","iopub.status.idle":"2025-02-25T20:10:07.41453Z","shell.execute_reply.started":"2025-02-25T20:10:05.358919Z","shell.execute_reply":"2025-02-25T20:10:07.413609Z"}},"outputs":[{"name":"stdout","text":"total 47M\n-rw-r--r-- 1 root root 6.7M Feb 25 20:01 X_test.parquet\n-rw-r--r-- 1 root root  20M Feb 25 20:01 X_train.parquet\n-rw-r--r-- 1 root root 6.7M Feb 25 20:01 X_val.parquet\n-rw-r--r-- 1 root root  205 Feb 25 20:09 dataset_info.joblib\n-rw-r--r-- 1 root root 2.7M Feb 25 20:01 y_test.csv\n-rw-r--r-- 1 root root 8.1M Feb 25 20:01 y_train.csv\n-rw-r--r-- 1 root root 2.7M Feb 25 20:01 y_val.csv\ntotal 8.0K\n-rw-r--r-- 1 root root  578 Feb 25 20:09 label_encoder.joblib\n-rw-r--r-- 1 root root 1007 Feb 25 20:09 standard_scaler.joblib\n","output_type":"stream"}],"execution_count":31},{"cell_type":"markdown","source":"# ================== LightGBM ==================","metadata":{}},{"cell_type":"code","source":"import lightgbm as lgb","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:21:45.369547Z","iopub.execute_input":"2025-02-25T20:21:45.369989Z","iopub.status.idle":"2025-02-25T20:21:45.374335Z","shell.execute_reply.started":"2025-02-25T20:21:45.36994Z","shell.execute_reply":"2025-02-25T20:21:45.37345Z"}},"outputs":[],"execution_count":43},{"cell_type":"code","source":"params = {\n    'learning_rate': 0.03,\n    'boosting_type': 'gbdt',\n    'objective': 'multiclass',\n    'metric': 'multi_logloss',\n    'max_depth': 7,\n    'num_class': 4,\n    'verbose': -1,\n    'device': 'gpu',\n    # 必须提前定义所有数据处理参数\n    'max_bin': 63,          # 必须在创建Dataset前定义\n    'bin_construct_sample_cnt': 200000,  # 优化GPU内存\n    'gpu_use_dp': True,      # 双精度模式\n    'seed': 1004,            # 固定随机种子\n    'feature_fraction': 0.7,\n    'bagging_freq': 5,\n    'lambda_l1': 0.1\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:21:46.612008Z","iopub.execute_input":"2025-02-25T20:21:46.612362Z","iopub.status.idle":"2025-02-25T20:21:46.617528Z","shell.execute_reply.started":"2025-02-25T20:21:46.61233Z","shell.execute_reply":"2025-02-25T20:21:46.616553Z"}},"outputs":[],"execution_count":44},{"cell_type":"markdown","source":"================== Передача params при создании Dataset ==================","metadata":{}},{"cell_type":"code","source":"# Обучающий набор (с весами)\nd_train = lgb.Dataset(\n    data=X_train.values,\n    label=y_train,\n    weight=sample_weights,  # 样本权重\n    params=params            # 关键！传入参数以锁定max_bin\n)\n\n# Валидационный набор (без весов)\nd_val = lgb.Dataset(\n    data=X_val.values,\n    label=y_val,\n    reference=d_train,       # 继承训练集的参数\n    params=params             # 显式传入相同参数\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:21:51.242505Z","iopub.execute_input":"2025-02-25T20:21:51.242857Z","iopub.status.idle":"2025-02-25T20:21:51.247859Z","shell.execute_reply.started":"2025-02-25T20:21:51.242823Z","shell.execute_reply":"2025-02-25T20:21:51.246988Z"}},"outputs":[],"execution_count":45},{"cell_type":"markdown","source":"================== 训练与验证 ==================","metadata":{}},{"cell_type":"code","source":"lightgbm_model = lgb.train(\n    params,\n    d_train,\n    num_boost_round=2000,\n    valid_sets=[d_val],\n    callbacks=[\n        lgb.early_stopping(stopping_rounds=20),\n        lgb.log_evaluation(20),\n        # 移除动态学习率回调（可能引发参数冲突）\n    ]\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:21:57.320462Z","iopub.execute_input":"2025-02-25T20:21:57.32081Z","iopub.status.idle":"2025-02-25T20:42:50.140757Z","shell.execute_reply.started":"2025-02-25T20:21:57.320778Z","shell.execute_reply":"2025-02-25T20:42:50.139831Z"}},"outputs":[{"name":"stderr","text":"1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n1 warning generated.\n","output_type":"stream"},{"name":"stdout","text":"Training until validation scores don't improve for 20 rounds\n[20]\tvalid_0's multi_logloss: 1.21427\n[40]\tvalid_0's multi_logloss: 1.15151\n[60]\tvalid_0's multi_logloss: 1.12203\n[80]\tvalid_0's multi_logloss: 1.10571\n[100]\tvalid_0's multi_logloss: 1.09713\n[120]\tvalid_0's multi_logloss: 1.09271\n[140]\tvalid_0's multi_logloss: 1.09022\n[160]\tvalid_0's multi_logloss: 1.08824\n[180]\tvalid_0's multi_logloss: 1.08706\n[200]\tvalid_0's multi_logloss: 1.08623\n[220]\tvalid_0's multi_logloss: 1.08563\n[240]\tvalid_0's multi_logloss: 1.08469\n[260]\tvalid_0's multi_logloss: 1.08385\n[280]\tvalid_0's multi_logloss: 1.08327\n[300]\tvalid_0's multi_logloss: 1.08288\n[320]\tvalid_0's multi_logloss: 1.08252\n[340]\tvalid_0's multi_logloss: 1.0823\n[360]\tvalid_0's multi_logloss: 1.08214\n[380]\tvalid_0's multi_logloss: 1.082\n[400]\tvalid_0's multi_logloss: 1.0818\n[420]\tvalid_0's multi_logloss: 1.08156\n[440]\tvalid_0's multi_logloss: 1.08132\n[460]\tvalid_0's multi_logloss: 1.08125\n[480]\tvalid_0's multi_logloss: 1.08118\n[500]\tvalid_0's multi_logloss: 1.08103\n[520]\tvalid_0's multi_logloss: 1.08092\n[540]\tvalid_0's multi_logloss: 1.08081\n[560]\tvalid_0's multi_logloss: 1.08071\n[580]\tvalid_0's multi_logloss: 1.08062\n[600]\tvalid_0's multi_logloss: 1.08052\n[620]\tvalid_0's multi_logloss: 1.08044\n[640]\tvalid_0's multi_logloss: 1.08032\n[660]\tvalid_0's multi_logloss: 1.08025\n[680]\tvalid_0's multi_logloss: 1.0802\n[700]\tvalid_0's multi_logloss: 1.08017\n[720]\tvalid_0's multi_logloss: 1.0801\n[740]\tvalid_0's multi_logloss: 1.08005\n[760]\tvalid_0's multi_logloss: 1.08001\n[780]\tvalid_0's multi_logloss: 1.07995\n[800]\tvalid_0's multi_logloss: 1.07991\n[820]\tvalid_0's multi_logloss: 1.07987\n[840]\tvalid_0's multi_logloss: 1.07983\n[860]\tvalid_0's multi_logloss: 1.0798\n[880]\tvalid_0's multi_logloss: 1.07976\n[900]\tvalid_0's multi_logloss: 1.07974\n[920]\tvalid_0's multi_logloss: 1.07968\n[940]\tvalid_0's multi_logloss: 1.07966\n[960]\tvalid_0's multi_logloss: 1.07963\n[980]\tvalid_0's multi_logloss: 1.07961\n[1000]\tvalid_0's multi_logloss: 1.07958\n[1020]\tvalid_0's multi_logloss: 1.07955\n[1040]\tvalid_0's multi_logloss: 1.07953\n[1060]\tvalid_0's multi_logloss: 1.07951\n[1080]\tvalid_0's multi_logloss: 1.07949\n[1100]\tvalid_0's multi_logloss: 1.07946\n[1120]\tvalid_0's multi_logloss: 1.07944\n[1140]\tvalid_0's multi_logloss: 1.07941\n[1160]\tvalid_0's multi_logloss: 1.0794\n[1180]\tvalid_0's multi_logloss: 1.07939\n[1200]\tvalid_0's multi_logloss: 1.07937\n[1220]\tvalid_0's multi_logloss: 1.07934\n[1240]\tvalid_0's multi_logloss: 1.07932\n[1260]\tvalid_0's multi_logloss: 1.07929\n[1280]\tvalid_0's multi_logloss: 1.07928\n[1300]\tvalid_0's multi_logloss: 1.07927\n[1320]\tvalid_0's multi_logloss: 1.07924\n[1340]\tvalid_0's multi_logloss: 1.07923\n[1360]\tvalid_0's multi_logloss: 1.07922\n[1380]\tvalid_0's multi_logloss: 1.07921\n[1400]\tvalid_0's multi_logloss: 1.07919\n[1420]\tvalid_0's multi_logloss: 1.07918\n[1440]\tvalid_0's multi_logloss: 1.07917\n[1460]\tvalid_0's multi_logloss: 1.07915\n[1480]\tvalid_0's multi_logloss: 1.07914\n[1500]\tvalid_0's multi_logloss: 1.07913\n[1520]\tvalid_0's multi_logloss: 1.07911\n[1540]\tvalid_0's multi_logloss: 1.0791\n[1560]\tvalid_0's multi_logloss: 1.07908\n[1580]\tvalid_0's multi_logloss: 1.07907\n[1600]\tvalid_0's multi_logloss: 1.07906\n[1620]\tvalid_0's multi_logloss: 1.07904\n[1640]\tvalid_0's multi_logloss: 1.07903\n[1660]\tvalid_0's multi_logloss: 1.07902\n[1680]\tvalid_0's multi_logloss: 1.079\n[1700]\tvalid_0's multi_logloss: 1.07899\n[1720]\tvalid_0's multi_logloss: 1.07897\n[1740]\tvalid_0's multi_logloss: 1.07895\n[1760]\tvalid_0's multi_logloss: 1.07894\n[1780]\tvalid_0's multi_logloss: 1.07893\n[1800]\tvalid_0's multi_logloss: 1.07891\n[1820]\tvalid_0's multi_logloss: 1.0789\n[1840]\tvalid_0's multi_logloss: 1.07889\n[1860]\tvalid_0's multi_logloss: 1.07888\n[1880]\tvalid_0's multi_logloss: 1.07888\n[1900]\tvalid_0's multi_logloss: 1.07887\n[1920]\tvalid_0's multi_logloss: 1.07886\n[1940]\tvalid_0's multi_logloss: 1.07886\n[1960]\tvalid_0's multi_logloss: 1.07885\n[1980]\tvalid_0's multi_logloss: 1.07885\n[2000]\tvalid_0's multi_logloss: 1.07883\nDid not meet early stopping. Best iteration is:\n[2000]\tvalid_0's multi_logloss: 1.07883\n","output_type":"stream"}],"execution_count":46},{"cell_type":"code","source":"from sklearn.metrics import classification_report\n\nval_pred = np.argmax(lightgbm_model.predict(X_val_scaled), axis=1)\nprint(\"验证集性能:\\n\", classification_report(y_val, val_pred, target_names=le.classes_))\n\ntest_pred = np.argmax(lightgbm_model.predict(X_test_scaled), axis=1)\nprint(\"测试集性能:\\n\", classification_report(y_test, test_pred, target_names=le.classes_))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:48:56.443047Z","iopub.execute_input":"2025-02-25T20:48:56.443415Z","iopub.status.idle":"2025-02-25T21:06:50.672035Z","shell.execute_reply.started":"2025-02-25T20:48:56.443368Z","shell.execute_reply":"2025-02-25T21:06:50.670926Z"}},"outputs":[{"name":"stdout","text":"验证集性能:\n                  precision    recall  f1-score   support\n\n         Normal       0.87      0.59      0.71    974253\nStartHesitation       0.11      0.48      0.18     60958\n           Turn       0.39      0.36      0.38    335757\n        Walking       0.10      0.42      0.16     41567\n\n       accuracy                           0.53   1412535\n      macro avg       0.37      0.47      0.36   1412535\n   weighted avg       0.70      0.53      0.59   1412535\n\n测试集性能:\n                  precision    recall  f1-score   support\n\n         Normal       0.87      0.59      0.71    974253\nStartHesitation       0.11      0.48      0.18     60958\n           Turn       0.39      0.36      0.38    335756\n        Walking       0.10      0.42      0.16     41568\n\n       accuracy                           0.53   1412535\n      macro avg       0.37      0.46      0.36   1412535\n   weighted avg       0.70      0.53      0.59   1412535\n\n","output_type":"stream"}],"execution_count":51},{"cell_type":"markdown","source":"# ================== XGBoost ==================","metadata":{}},{"cell_type":"code","source":"import xgboost as xgb\nfrom sklearn.metrics import classification_report","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T21:06:50.674048Z","iopub.execute_input":"2025-02-25T21:06:50.674492Z","iopub.status.idle":"2025-02-25T21:06:50.955138Z","shell.execute_reply.started":"2025-02-25T21:06:50.674441Z","shell.execute_reply":"2025-02-25T21:06:50.954216Z"}},"outputs":[],"execution_count":52},{"cell_type":"markdown","source":"================== 1. Загрузка данных и предварительная обработка ==================","metadata":{}},{"cell_type":"code","source":"# Загрузка признаков\nX_train = pd.read_parquet('/kaggle/working/datasets/X_train.parquet')\nX_val = pd.read_parquet('/kaggle/working/datasets/X_val.parquet')\nX_test = pd.read_parquet('/kaggle/working/datasets/X_test.parquet')\n\n# Загрузка меток\ny_train = pd.read_csv('/kaggle/working/datasets/y_train.csv')['label']\ny_val = pd.read_csv('/kaggle/working/datasets/y_val.csv')['label']\ny_test = pd.read_csv('/kaggle/working/datasets/y_test.csv')['label']\n\n# Загрузка объекта предварительной обработки\nscaler = joblib.load('/kaggle/working/preprocessors/standard_scaler.joblib')\nle = joblib.load('/kaggle/working/preprocessors/label_encoder.joblib')\n\n# Применение стандартизации\nX_train_scaled = scaler.transform(X_train)\nX_val_scaled = scaler.transform(X_val)\nX_test_scaled = scaler.transform(X_test)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T21:06:50.956225Z","iopub.execute_input":"2025-02-25T21:06:50.956498Z","iopub.status.idle":"2025-02-25T21:06:51.980396Z","shell.execute_reply.started":"2025-02-25T21:06:50.956471Z","shell.execute_reply":"2025-02-25T21:06:51.979358Z"}},"outputs":[],"execution_count":53},{"cell_type":"markdown","source":"================== 2. Расчет весов выборки ==================","metadata":{}},{"cell_type":"code","source":"# Расчет весов на основе распределения меток обучающего набора\nclass_weights = compute_class_weight(\n    class_weight='balanced',\n    classes=np.unique(y_train),\n    y=y_train\n)\nsample_weights = np.array([class_weights[y] for y in y_train])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T21:06:51.982262Z","iopub.execute_input":"2025-02-25T21:06:51.982547Z","iopub.status.idle":"2025-02-25T21:06:53.818397Z","shell.execute_reply.started":"2025-02-25T21:06:51.982519Z","shell.execute_reply":"2025-02-25T21:06:53.817639Z"}},"outputs":[],"execution_count":54},{"cell_type":"markdown","source":"================== 3.Настройка параметров XGBoost ==================","metadata":{}},{"cell_type":"code","source":"params = {\n    'learning_rate': 0.5,\n    'max_depth': 6,\n    'min_child_weight': 10,\n    'random_state': 30,\n    'objective': 'multi:softmax',\n    'num_class': 4,\n    'tree_method': 'hist',          \n    'device': 'cuda:0',             \n    'eval_metric': 'mlogloss',\n    # 可选优化参数\n    'subsample': 0.8,               # Случайная выборка образцов\n    'colsample_bytree': 0.8         # Случайная выборка признаков\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T21:06:53.81962Z","iopub.execute_input":"2025-02-25T21:06:53.820279Z","iopub.status.idle":"2025-02-25T21:06:53.82479Z","shell.execute_reply.started":"2025-02-25T21:06:53.820246Z","shell.execute_reply":"2025-02-25T21:06:53.823763Z"}},"outputs":[],"execution_count":55},{"cell_type":"markdown","source":"================== 4.Обучение модели (с ранней остановкой) ==================","metadata":{}},{"cell_type":"code","source":"dtrain = xgb.DMatrix(X_train_scaled, label=y_train, weight=sample_weights)\ndval = xgb.DMatrix(X_val_scaled, label=y_val)\n\nxgboost_model = xgb.train(\n    params,\n    dtrain,\n    num_boost_round=2000,\n    evals=[(dtrain, 'train'), (dval, 'val')],\n    early_stopping_rounds=50,\n    verbose_eval=50\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T21:06:53.825988Z","iopub.execute_input":"2025-02-25T21:06:53.826378Z","iopub.status.idle":"2025-02-25T21:08:13.762016Z","shell.execute_reply.started":"2025-02-25T21:06:53.82634Z","shell.execute_reply":"2025-02-25T21:08:13.761038Z"}},"outputs":[{"name":"stdout","text":"[0]\ttrain-mlogloss:1.29055\tval-mlogloss:1.23272\n[50]\ttrain-mlogloss:1.13779\tval-mlogloss:1.05782\n[100]\ttrain-mlogloss:1.13148\tval-mlogloss:1.05559\n[150]\ttrain-mlogloss:1.12697\tval-mlogloss:1.05426\n[200]\ttrain-mlogloss:1.12313\tval-mlogloss:1.05400\n[250]\ttrain-mlogloss:1.11982\tval-mlogloss:1.05339\n[300]\ttrain-mlogloss:1.11688\tval-mlogloss:1.05316\n[350]\ttrain-mlogloss:1.11421\tval-mlogloss:1.05276\n[400]\ttrain-mlogloss:1.11171\tval-mlogloss:1.05268\n[450]\ttrain-mlogloss:1.10934\tval-mlogloss:1.05240\n[500]\ttrain-mlogloss:1.10707\tval-mlogloss:1.05199\n[550]\ttrain-mlogloss:1.10493\tval-mlogloss:1.05211\n[600]\ttrain-mlogloss:1.10299\tval-mlogloss:1.05177\n[650]\ttrain-mlogloss:1.10118\tval-mlogloss:1.05157\n[688]\ttrain-mlogloss:1.09983\tval-mlogloss:1.05167\n","output_type":"stream"}],"execution_count":56},{"cell_type":"markdown","source":"================== 5. Оценка на валидационном наборе ==================","metadata":{}},{"cell_type":"code","source":"y_val_pred = xgboost_model.predict(dval)\nprint(\"\\n==== 验证集分类报告 (XGBoost) ====\")\nprint(classification_report(\n    y_val, y_val_pred,\n    target_names=le.classes_,\n    digits=4\n))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T21:08:13.768409Z","iopub.execute_input":"2025-02-25T21:08:13.768706Z","iopub.status.idle":"2025-02-25T21:08:16.262508Z","shell.execute_reply.started":"2025-02-25T21:08:13.768678Z","shell.execute_reply":"2025-02-25T21:08:16.261555Z"}},"outputs":[{"name":"stdout","text":"\n==== 验证集分类报告 (XGBoost) ====\n                 precision    recall  f1-score   support\n\n         Normal     0.8749    0.6061    0.7161    974253\nStartHesitation     0.1200    0.4770    0.1917     60958\n           Turn     0.4050    0.3861    0.3953    335757\n        Walking     0.0971    0.4091    0.1570     41567\n\n       accuracy                         0.5424   1412535\n      macro avg     0.3742    0.4696    0.3650   1412535\n   weighted avg     0.7077    0.5424    0.6008   1412535\n\n","output_type":"stream"}],"execution_count":57},{"cell_type":"markdown","source":"================== 6. Сохранение модели ==================","metadata":{}},{"cell_type":"code","source":"model.save_model('/kaggle/working/xgboost_model.json')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-24T21:46:47.588006Z","iopub.execute_input":"2025-02-24T21:46:47.588713Z","iopub.status.idle":"2025-02-24T21:46:47.633457Z","shell.execute_reply.started":"2025-02-24T21:46:47.588683Z","shell.execute_reply":"2025-02-24T21:46:47.632569Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"================== 7. Сравнение с результатами LightGBM ==================","metadata":{}},{"cell_type":"code","source":"#","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# ================== CNN ==================","metadata":{}},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import Conv1D, MaxPooling1D, Flatten, Dense, Dropout\nfrom tensorflow.keras.utils import to_categorical\nfrom sklearn.metrics import classification_report","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:02:45.572633Z","iopub.execute_input":"2025-02-25T20:02:45.573038Z","iopub.status.idle":"2025-02-25T20:02:59.184893Z","shell.execute_reply.started":"2025-02-25T20:02:45.573Z","shell.execute_reply":"2025-02-25T20:02:59.183917Z"}},"outputs":[],"execution_count":19},{"cell_type":"markdown","source":"================== 1. Загрузка данных и предварительная обработка ==================","metadata":{}},{"cell_type":"code","source":"# 加载数据\nX_train = pd.read_parquet('/kaggle/working/datasets/X_train.parquet')\ny_train = pd.read_csv('/kaggle/working/datasets/y_train.csv')['label']\nX_val = pd.read_parquet('/kaggle/working/datasets/X_val.parquet')\ny_val = pd.read_csv('/kaggle/working/datasets/y_val.csv')['label']\nX_test = pd.read_parquet('/kaggle/working/datasets/X_test.parquet')\ny_test = pd.read_csv('/kaggle/working/datasets/y_test.csv')['label']","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:11:38.770654Z","iopub.execute_input":"2025-02-25T20:11:38.771618Z","iopub.status.idle":"2025-02-25T20:11:39.232749Z","shell.execute_reply.started":"2025-02-25T20:11:38.771579Z","shell.execute_reply":"2025-02-25T20:11:39.231925Z"}},"outputs":[],"execution_count":32},{"cell_type":"code","source":"# Загрузка информации о предварительной обработке\ndataset_info = joblib.load('/kaggle/working/datasets/dataset_info.joblib')\nle = joblib.load('/kaggle/working/preprocessors/label_encoder.joblib')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:11:40.989988Z","iopub.execute_input":"2025-02-25T20:11:40.990804Z","iopub.status.idle":"2025-02-25T20:11:40.995782Z","shell.execute_reply.started":"2025-02-25T20:11:40.990767Z","shell.execute_reply":"2025-02-25T20:11:40.995107Z"}},"outputs":[],"execution_count":33},{"cell_type":"code","source":"# Адаптация формы данных (добавление измерения временного шага)\nX_train_3d = X_train.values.reshape(-1, 1, 3)  # (samples, timesteps=1, features=3)\nX_val_3d = X_val.values.reshape(-1, 1, 3)\nX_test_3d = X_test.values.reshape(-1, 1, 3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:11:42.389499Z","iopub.execute_input":"2025-02-25T20:11:42.38985Z","iopub.status.idle":"2025-02-25T20:11:42.394754Z","shell.execute_reply.started":"2025-02-25T20:11:42.389817Z","shell.execute_reply":"2025-02-25T20:11:42.393925Z"}},"outputs":[],"execution_count":34},{"cell_type":"code","source":"# Преобразование меток в one-hot\ny_train_cat = to_categorical(y_train)\ny_val_cat = to_categorical(y_val)\ny_test_cat = to_categorical(y_test)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:11:43.700363Z","iopub.execute_input":"2025-02-25T20:11:43.701325Z","iopub.status.idle":"2025-02-25T20:11:43.951661Z","shell.execute_reply.started":"2025-02-25T20:11:43.701287Z","shell.execute_reply":"2025-02-25T20:11:43.950661Z"}},"outputs":[],"execution_count":35},{"cell_type":"code","source":"# Расчет весов категорий (адаптация к формату Keras)\ndataset_info = joblib.load('/kaggle/working/datasets/dataset_info.joblib')\nclass_weights = np.array(dataset_info['class_weights'])  # 转换为numpy数组\nclass_weights_dict = {i: w for i, w in enumerate(class_weights)}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:11:46.564685Z","iopub.execute_input":"2025-02-25T20:11:46.565384Z","iopub.status.idle":"2025-02-25T20:11:46.570433Z","shell.execute_reply.started":"2025-02-25T20:11:46.565348Z","shell.execute_reply":"2025-02-25T20:11:46.569593Z"}},"outputs":[],"execution_count":36},{"cell_type":"markdown","source":"================== 2.Простая архитектура модели 1D-CNN ==================","metadata":{}},{"cell_type":"code","source":"# Построение модели CNN\ndef build_cnn_model():\n    model = Sequential([\n        Conv1D(64, kernel_size=1, activation='relu', \n              input_shape=(1, 3), name='conv1d_first'),\n        MaxPooling1D(pool_size=1, name='maxpool'),\n        Flatten(name='flatten'),\n        Dense(128, activation='relu', name='fc1'),\n        Dropout(0.5, name='dropout'),\n        Dense(64, activation='relu', name='fc2'),\n        Dense(dataset_info['num_classes'], activation='softmax', name='output')\n    ])\n    \n    model.compile(\n        optimizer=tf.keras.optimizers.Adam(learning_rate=0.001),\n        loss='categorical_crossentropy',\n        metrics=['accuracy', \n                tf.keras.metrics.Precision(name='precision'),\n                tf.keras.metrics.Recall(name='recall')]\n    )\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:12:33.149854Z","iopub.execute_input":"2025-02-25T20:12:33.150741Z","iopub.status.idle":"2025-02-25T20:12:33.157046Z","shell.execute_reply.started":"2025-02-25T20:12:33.150708Z","shell.execute_reply":"2025-02-25T20:12:33.156061Z"}},"outputs":[],"execution_count":37},{"cell_type":"code","source":"# Инициализация модели\ncnn_model = build_cnn_model()\ncnn_model.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:12:34.75831Z","iopub.execute_input":"2025-02-25T20:12:34.758657Z","iopub.status.idle":"2025-02-25T20:12:35.655455Z","shell.execute_reply.started":"2025-02-25T20:12:34.758624Z","shell.execute_reply":"2025-02-25T20:12:35.654577Z"}},"outputs":[{"name":"stderr","text":"/opt/conda/lib/python3.10/site-packages/keras/src/layers/convolutional/base_conv.py:107: UserWarning: Do not pass an `input_shape`/`input_dim` argument to a layer. When using Sequential models, prefer using an `Input(shape)` object as the first layer in the model instead.\n  super().__init__(activity_regularizer=activity_regularizer, **kwargs)\n","output_type":"stream"},{"output_type":"display_data","data":{"text/plain":"\u001b[1mModel: \"sequential\"\u001b[0m\n","text/html":"<pre style=\"white-space:pre;overflow-x:auto;line-height:normal;font-family:Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace\"><span style=\"font-weight: bold\">Model: \"sequential\"</span>\n</pre>\n"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"┏━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━┓\n┃\u001b[1m \u001b[0m\u001b[1mLayer (type)                   \u001b[0m\u001b[1m \u001b[0m┃\u001b[1m \u001b[0m\u001b[1mOutput Shape          \u001b[0m\u001b[1m \u001b[0m┃\u001b[1m \u001b[0m\u001b[1m      Param #\u001b[0m\u001b[1m \u001b[0m┃\n┡━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━┩\n│ conv1d_first (\u001b[38;5;33mConv1D\u001b[0m)           │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m1\u001b[0m, \u001b[38;5;34m64\u001b[0m)          │           \u001b[38;5;34m256\u001b[0m │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ maxpool (\u001b[38;5;33mMaxPooling1D\u001b[0m)          │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m1\u001b[0m, \u001b[38;5;34m64\u001b[0m)          │             \u001b[38;5;34m0\u001b[0m │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ flatten (\u001b[38;5;33mFlatten\u001b[0m)               │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m64\u001b[0m)             │             \u001b[38;5;34m0\u001b[0m │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ fc1 (\u001b[38;5;33mDense\u001b[0m)                     │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m128\u001b[0m)            │         \u001b[38;5;34m8,320\u001b[0m │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ dropout (\u001b[38;5;33mDropout\u001b[0m)               │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m128\u001b[0m)            │             \u001b[38;5;34m0\u001b[0m │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ fc2 (\u001b[38;5;33mDense\u001b[0m)                     │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m64\u001b[0m)             │         \u001b[38;5;34m8,256\u001b[0m │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ output (\u001b[38;5;33mDense\u001b[0m)                  │ (\u001b[38;5;45mNone\u001b[0m, \u001b[38;5;34m4\u001b[0m)              │           \u001b[38;5;34m260\u001b[0m │\n└─────────────────────────────────┴────────────────────────┴───────────────┘\n","text/html":"<pre style=\"white-space:pre;overflow-x:auto;line-height:normal;font-family:Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace\">┏━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━━━━━━━━━━┳━━━━━━━━━━━━━━━┓\n┃<span style=\"font-weight: bold\"> Layer (type)                    </span>┃<span style=\"font-weight: bold\"> Output Shape           </span>┃<span style=\"font-weight: bold\">       Param # </span>┃\n┡━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━━━━━━━━━━╇━━━━━━━━━━━━━━━┩\n│ conv1d_first (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Conv1D</span>)           │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">1</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">64</span>)          │           <span style=\"color: #00af00; text-decoration-color: #00af00\">256</span> │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ maxpool (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">MaxPooling1D</span>)          │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">1</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">64</span>)          │             <span style=\"color: #00af00; text-decoration-color: #00af00\">0</span> │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ flatten (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Flatten</span>)               │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">64</span>)             │             <span style=\"color: #00af00; text-decoration-color: #00af00\">0</span> │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ fc1 (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Dense</span>)                     │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">128</span>)            │         <span style=\"color: #00af00; text-decoration-color: #00af00\">8,320</span> │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ dropout (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Dropout</span>)               │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">128</span>)            │             <span style=\"color: #00af00; text-decoration-color: #00af00\">0</span> │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ fc2 (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Dense</span>)                     │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">64</span>)             │         <span style=\"color: #00af00; text-decoration-color: #00af00\">8,256</span> │\n├─────────────────────────────────┼────────────────────────┼───────────────┤\n│ output (<span style=\"color: #0087ff; text-decoration-color: #0087ff\">Dense</span>)                  │ (<span style=\"color: #00d7ff; text-decoration-color: #00d7ff\">None</span>, <span style=\"color: #00af00; text-decoration-color: #00af00\">4</span>)              │           <span style=\"color: #00af00; text-decoration-color: #00af00\">260</span> │\n└─────────────────────────────────┴────────────────────────┴───────────────┘\n</pre>\n"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"\u001b[1m Total params: \u001b[0m\u001b[38;5;34m17,092\u001b[0m (66.77 KB)\n","text/html":"<pre style=\"white-space:pre;overflow-x:auto;line-height:normal;font-family:Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace\"><span style=\"font-weight: bold\"> Total params: </span><span style=\"color: #00af00; text-decoration-color: #00af00\">17,092</span> (66.77 KB)\n</pre>\n"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"\u001b[1m Trainable params: \u001b[0m\u001b[38;5;34m17,092\u001b[0m (66.77 KB)\n","text/html":"<pre style=\"white-space:pre;overflow-x:auto;line-height:normal;font-family:Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace\"><span style=\"font-weight: bold\"> Trainable params: </span><span style=\"color: #00af00; text-decoration-color: #00af00\">17,092</span> (66.77 KB)\n</pre>\n"},"metadata":{}},{"output_type":"display_data","data":{"text/plain":"\u001b[1m Non-trainable params: \u001b[0m\u001b[38;5;34m0\u001b[0m (0.00 B)\n","text/html":"<pre style=\"white-space:pre;overflow-x:auto;line-height:normal;font-family:Menlo,'DejaVu Sans Mono',consolas,'Courier New',monospace\"><span style=\"font-weight: bold\"> Non-trainable params: </span><span style=\"color: #00af00; text-decoration-color: #00af00\">0</span> (0.00 B)\n</pre>\n"},"metadata":{}}],"execution_count":38},{"cell_type":"markdown","source":"================== 3. Конфигурация обучения ==================","metadata":{}},{"cell_type":"code","source":"# Улучшенная конфигурация обучения\ncallbacks = [\n    tf.keras.callbacks.EarlyStopping(patience=7, restore_best_weights=True),\n    tf.keras.callbacks.ReduceLROnPlateau(factor=0.5, patience=3)\n]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:12:44.771706Z","iopub.execute_input":"2025-02-25T20:12:44.772081Z","iopub.status.idle":"2025-02-25T20:12:44.776833Z","shell.execute_reply.started":"2025-02-25T20:12:44.772046Z","shell.execute_reply":"2025-02-25T20:12:44.77582Z"}},"outputs":[],"execution_count":39},{"cell_type":"markdown","source":"================== 4. Выполнение обучения ==================","metadata":{}},{"cell_type":"code","source":"# 开始训练\nhistory = cnn_model.fit(\n    X_train_3d, y_train_cat,\n    validation_data=(X_val_3d, y_val_cat),\n    epochs=50,\n    batch_size=1024,\n    class_weight=class_weights_dict,\n    callbacks=callbacks,\n    verbose=1\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:12:47.599714Z","iopub.execute_input":"2025-02-25T20:12:47.600369Z","iopub.status.idle":"2025-02-25T20:15:56.404088Z","shell.execute_reply.started":"2025-02-25T20:12:47.600333Z","shell.execute_reply":"2025-02-25T20:15:56.403307Z"}},"outputs":[{"name":"stdout","text":"Epoch 1/50\n","output_type":"stream"},{"name":"stderr","text":"WARNING: All log messages before absl::InitializeLog() is called are written to STDERR\nI0000 00:00:1740514371.362813     201 service.cc:145] XLA service 0x5894ecf8c6a0 initialized for platform CUDA (this does not guarantee that XLA will be used). Devices:\nI0000 00:00:1740514371.362860     201 service.cc:153]   StreamExecutor device (0): Tesla T4, Compute Capability 7.5\nI0000 00:00:1740514371.362863     201 service.cc:153]   StreamExecutor device (1): Tesla T4, Compute Capability 7.5\n","output_type":"stream"},{"name":"stdout","text":"\u001b[1m  88/4139\u001b[0m \u001b[37m━━━━━━━━━━━━━━━━━━━━\u001b[0m \u001b[1m7s\u001b[0m 2ms/step - accuracy: 0.4219 - loss: 1.3371 - precision: 0.7224 - recall: 0.0785          ","output_type":"stream"},{"name":"stderr","text":"I0000 00:00:1740514375.098689     201 device_compiler.h:188] Compiled cluster using XLA!  This line is logged at most once for the lifetime of the process.\n","output_type":"stream"},{"name":"stdout","text":"\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m18s\u001b[0m 3ms/step - accuracy: 0.5164 - loss: 1.2065 - precision: 0.8653 - recall: 0.2563 - val_accuracy: 0.5656 - val_loss: 1.0670 - val_precision: 0.8206 - val_recall: 0.3145 - learning_rate: 0.0010\nEpoch 2/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5541 - loss: 1.1633 - precision: 0.8262 - recall: 0.3032 - val_accuracy: 0.5614 - val_loss: 1.0677 - val_precision: 0.8173 - val_recall: 0.3077 - learning_rate: 0.0010\nEpoch 3/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5559 - loss: 1.1589 - precision: 0.8197 - recall: 0.3068 - val_accuracy: 0.5589 - val_loss: 1.0661 - val_precision: 0.8147 - val_recall: 0.3105 - learning_rate: 0.0010\nEpoch 4/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5559 - loss: 1.1558 - precision: 0.8174 - recall: 0.3107 - val_accuracy: 0.5688 - val_loss: 1.0528 - val_precision: 0.8272 - val_recall: 0.3113 - learning_rate: 0.0010\nEpoch 5/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5587 - loss: 1.1541 - precision: 0.8115 - recall: 0.3126 - val_accuracy: 0.5674 - val_loss: 1.0694 - val_precision: 0.8119 - val_recall: 0.3047 - learning_rate: 0.0010\nEpoch 6/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5579 - loss: 1.1534 - precision: 0.8126 - recall: 0.3140 - val_accuracy: 0.5678 - val_loss: 1.0569 - val_precision: 0.8176 - val_recall: 0.3218 - learning_rate: 0.0010\nEpoch 7/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5584 - loss: 1.1506 - precision: 0.8086 - recall: 0.3156 - val_accuracy: 0.5675 - val_loss: 1.0598 - val_precision: 0.8151 - val_recall: 0.3110 - learning_rate: 0.0010\nEpoch 8/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5593 - loss: 1.1511 - precision: 0.8080 - recall: 0.3180 - val_accuracy: 0.5658 - val_loss: 1.0550 - val_precision: 0.8200 - val_recall: 0.3111 - learning_rate: 5.0000e-04\nEpoch 9/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5592 - loss: 1.1485 - precision: 0.8086 - recall: 0.3179 - val_accuracy: 0.5654 - val_loss: 1.0507 - val_precision: 0.8181 - val_recall: 0.3133 - learning_rate: 5.0000e-04\nEpoch 10/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m10s\u001b[0m 2ms/step - accuracy: 0.5584 - loss: 1.1491 - precision: 0.8073 - recall: 0.3181 - val_accuracy: 0.5660 - val_loss: 1.0567 - val_precision: 0.8080 - val_recall: 0.3151 - learning_rate: 5.0000e-04\nEpoch 11/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5601 - loss: 1.1452 - precision: 0.8069 - recall: 0.3196 - val_accuracy: 0.5606 - val_loss: 1.0637 - val_precision: 0.7983 - val_recall: 0.3152 - learning_rate: 5.0000e-04\nEpoch 12/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5585 - loss: 1.1484 - precision: 0.8051 - recall: 0.3209 - val_accuracy: 0.5770 - val_loss: 1.0482 - val_precision: 0.8114 - val_recall: 0.3187 - learning_rate: 5.0000e-04\nEpoch 13/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m10s\u001b[0m 2ms/step - accuracy: 0.5610 - loss: 1.1453 - precision: 0.8049 - recall: 0.3198 - val_accuracy: 0.5577 - val_loss: 1.0667 - val_precision: 0.8219 - val_recall: 0.3074 - learning_rate: 5.0000e-04\nEpoch 14/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5592 - loss: 1.1464 - precision: 0.8049 - recall: 0.3192 - val_accuracy: 0.5621 - val_loss: 1.0670 - val_precision: 0.8191 - val_recall: 0.3049 - learning_rate: 5.0000e-04\nEpoch 15/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5605 - loss: 1.1460 - precision: 0.8036 - recall: 0.3200 - val_accuracy: 0.5615 - val_loss: 1.0562 - val_precision: 0.8070 - val_recall: 0.3127 - learning_rate: 5.0000e-04\nEpoch 16/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5607 - loss: 1.1450 - precision: 0.8010 - recall: 0.3212 - val_accuracy: 0.5605 - val_loss: 1.0665 - val_precision: 0.8239 - val_recall: 0.2994 - learning_rate: 2.5000e-04\nEpoch 17/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5602 - loss: 1.1457 - precision: 0.8043 - recall: 0.3197 - val_accuracy: 0.5625 - val_loss: 1.0619 - val_precision: 0.8081 - val_recall: 0.3042 - learning_rate: 2.5000e-04\nEpoch 18/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5617 - loss: 1.1424 - precision: 0.8031 - recall: 0.3210 - val_accuracy: 0.5606 - val_loss: 1.0651 - val_precision: 0.8138 - val_recall: 0.3013 - learning_rate: 2.5000e-04\nEpoch 19/50\n\u001b[1m4139/4139\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m9s\u001b[0m 2ms/step - accuracy: 0.5618 - loss: 1.1431 - precision: 0.8064 - recall: 0.3200 - val_accuracy: 0.5608 - val_loss: 1.0641 - val_precision: 0.8078 - val_recall: 0.3055 - learning_rate: 1.2500e-04\n","output_type":"stream"}],"execution_count":40},{"cell_type":"code","source":"# Оценка на тестовом наборе\ntest_results = cnn_model.evaluate(X_test_3d, y_test_cat, verbose=0)\nprint(\"\\nTest Metrics:\")\nprint(f\"Loss: {test_results[0]:.4f}\")\nprint(f\"Accuracy: {test_results[1]:.4f}\")\nprint(f\"Precision: {test_results[2]:.4f}\")\nprint(f\"Recall: {test_results[3]:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:17:24.263079Z","iopub.execute_input":"2025-02-25T20:17:24.263875Z","iopub.status.idle":"2025-02-25T20:18:18.503455Z","shell.execute_reply.started":"2025-02-25T20:17:24.263824Z","shell.execute_reply":"2025-02-25T20:18:18.502385Z"}},"outputs":[{"name":"stdout","text":"\nTest Metrics:\nLoss: 1.0497\nAccuracy: 0.5758\nPrecision: 0.8106\nRecall: 0.3173\n","output_type":"stream"}],"execution_count":41},{"cell_type":"code","source":"# Генерация отчета о классификации\ny_pred = np.argmax(cnn_model.predict(X_test_3d), axis=1)\nprint(\"\\nClassification Report:\")\nprint(classification_report(y_test, y_pred, target_names=le.classes_))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T20:18:29.563789Z","iopub.execute_input":"2025-02-25T20:18:29.564595Z","iopub.status.idle":"2025-02-25T20:19:49.762503Z","shell.execute_reply.started":"2025-02-25T20:18:29.564561Z","shell.execute_reply":"2025-02-25T20:19:49.76154Z"}},"outputs":[{"name":"stdout","text":"\u001b[1m44142/44142\u001b[0m \u001b[32m━━━━━━━━━━━━━━━━━━━━\u001b[0m\u001b[37m\u001b[0m \u001b[1m50s\u001b[0m 1ms/step\n\nClassification Report:\n                 precision    recall  f1-score   support\n\n         Normal       0.87      0.65      0.74    974253\nStartHesitation       0.12      0.49      0.20     60958\n           Turn       0.43      0.39      0.41    335756\n        Walking       0.12      0.38      0.18     41568\n\n       accuracy                           0.58   1412535\n      macro avg       0.38      0.48      0.38   1412535\n   weighted avg       0.71      0.58      0.62   1412535\n\n","output_type":"stream"}],"execution_count":42},{"cell_type":"code","source":"# Сохранение полной модели\ncnn_model.save('/kaggle/working/cnn_parkinsons_model.h5')","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 此部分是1DCNN文献复现代码，使用resnet模型和5D交叉验证，可参考部分优化的思路\n### Эта часть представляет собой код воспроизведения статьи по 1DCNN, использует модель resnet и 5D кросс-валидацию, можно использовать для частичной оптимизации идей\n#### Еще не удалось успешно запустить.","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport gc\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedGroupKFold","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-25T19:57:21.957025Z","iopub.execute_input":"2025-02-25T19:57:21.957672Z","iopub.status.idle":"2025-02-25T19:57:21.994608Z","shell.execute_reply.started":"2025-02-25T19:57:21.957636Z","shell.execute_reply":"2025-02-25T19:57:21.993519Z"}},"outputs":[{"traceback":["\u001b[0;31m---------------------------------------------------------------------------\u001b[0m","\u001b[0;31mModuleNotFoundError\u001b[0m                       Traceback (most recent call last)","Cell \u001b[0;32mIn[7], line 10\u001b[0m\n\u001b[1;32m      8\u001b[0m \u001b[38;5;28;01mfrom\u001b[39;00m \u001b[38;5;21;01mtorch\u001b[39;00m\u001b[38;5;21;01m.\u001b[39;00m\u001b[38;5;21;01mutils\u001b[39;00m\u001b[38;5;21;01m.\u001b[39;00m\u001b[38;5;21;01mdata\u001b[39;00m \u001b[38;5;28;01mimport\u001b[39;00m Dataset, DataLoader\n\u001b[1;32m      9\u001b[0m \u001b[38;5;28;01mfrom\u001b[39;00m \u001b[38;5;21;01msklearn\u001b[39;00m\u001b[38;5;21;01m.\u001b[39;00m\u001b[38;5;21;01mmodel_selection\u001b[39;00m \u001b[38;5;28;01mimport\u001b[39;00m train_test_split\n\u001b[0;32m---> 10\u001b[0m \u001b[38;5;28;01mfrom\u001b[39;00m \u001b[38;5;21;01miterstrat\u001b[39;00m\u001b[38;5;21;01m.\u001b[39;00m\u001b[38;5;21;01mml_stratifiers\u001b[39;00m \u001b[38;5;28;01mimport\u001b[39;00m MultilabelStratifiedGroupKFold\n","\u001b[0;31mModuleNotFoundError\u001b[0m: No module named 'iterstrat'"],"ename":"ModuleNotFoundError","evalue":"No module named 'iterstrat'","output_type":"error"}],"execution_count":7},{"cell_type":"code","source":"# 配置参数\nclass Config:\n    window_size = 1000\n    features = ['V', 'ML', 'AP']\n    target_cols = ['StartHesitation', 'Turn', 'Walking']\n    batch_size = 1024\n    num_workers = 4\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    num_epochs = 1\n    max_samples = 5_000_000\n    n_folds = 5\n    lr = 0.001","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 数据加载和预处理（修改部分）\ndef load_and_preprocess():\n    # 保持原有数据加载逻辑\n    DATA_ROOT_TDCSFOG = '/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/train/tdcsfog/'\n    tdcsfog = pd.concat([pd.read_csv(os.path.join(root, name)).assign(file=name.split('.')[0]) \n                       for root, _, files in os.walk(DATA_ROOT_TDCSFOG) \n                       for name in files], axis=0)\n    tdcsfog = reduce_memory_usage(tdcsfog)\n    \n    tdcsfog_metadata = pd.read_csv(\"/kaggle/input/tlvmc-parkinsons-freezing-gait-prediction/tdcsfog_metadata.csv\")\n    tdcsfog_m = tdcsfog_metadata.merge(tdcsfog, how='inner', left_on='Id', right_on='file').drop('file', axis=1)\n    \n    # 窗口生成\n    samples, labels, groups = [], [], []\n    for Id, group in tdcsfog_m.groupby('Id'):\n        group = group.sort_values('Time')\n        data = group[Config.features].values\n        targets = group[Config.target_cols].values\n        \n        for i in range(len(data) // Config.window_size):\n            start = i * Config.window_size\n            end = start + Config.window_size\n            window_data = data[start:end]\n            window_label = (targets[start:end].sum(axis=0) > 0).astype(np.float32)\n            \n            samples.append(window_data)\n            labels.append(window_label)\n            groups.append(Id)\n    \n    return np.array(samples), np.array(labels), np.array(groups)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# PyTorch数据集类\nclass FOGDataset(Dataset):\n    def __init__(self, X, y):\n        self.X = X\n        self.y = y\n        \n    def __len__(self):\n        return len(self.X)\n    \n    def __getitem__(self, idx):\n        return (\n            torch.FloatTensor(self.X[idx].transpose(1, 0)),  # (3, 1000)\n            torch.FloatTensor(self.y[idx])                   # (3,)\n        )","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 1D残差块\nclass ResidualBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, stride, kernel_size=17):\n        super().__init__()\n        self.conv1 = nn.Conv1d(in_channels, out_channels, \n                              kernel_size=kernel_size, \n                              stride=stride,\n                              padding=kernel_size//2)\n        self.bn1 = nn.BatchNorm1d(out_channels)\n        self.conv2 = nn.Conv1d(out_channels, out_channels, \n                              kernel_size=kernel_size, \n                              stride=1,\n                              padding=kernel_size//2)\n        self.bn2 = nn.BatchNorm1d(out_channels)\n        \n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv1d(in_channels, out_channels, \n                         kernel_size=1, \n                         stride=stride),\n                nn.BatchNorm1d(out_channels)\n            )\n    \n    def forward(self, x):\n        residual = self.shortcut(x)\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = nn.ReLU()(x)\n        x = self.conv2(x)\n        x = self.bn2(x)\n        x += residual\n        return nn.ReLU()(x)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ResNet模型\nclass ResNet1D(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.layers = nn.Sequential(\n            ResidualBlock(3, 32, stride=2),\n            ResidualBlock(32, 64, stride=2),\n            ResidualBlock(64, 128, stride=2),\n            ResidualBlock(128, 256, stride=5, kernel_size=5),\n            ResidualBlock(256, 384, stride=5, kernel_size=5),\n            ResidualBlock(384, 512, stride=1)\n        )\n        self.avgpool = nn.AdaptiveAvgPool1d(1)\n        self.fc = nn.Linear(512, 3)\n        \n    def forward(self, x):\n        x = self.layers(x)           # (batch, 512, 5)\n        x = self.avgpool(x)          # (batch, 512, 1)\n        x = x.view(x.size(0), -1)    # (batch, 512)\n        return self.fc(x)            # (batch, 3)\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 训练函数\ndef train_fold(fold, X, y, groups):\n    sgkf = MultilabelStratifiedGroupKFold(n_splits=Config.n_folds, shuffle=True)\n    \n    for fold, (train_idx, val_idx) in enumerate(sgkf.split(X, y, groups)):\n        print(f\"Training Fold {fold+1}\")\n        \n        # 数据加载器\n        train_dataset = FOGDataset(X[train_idx], y[train_idx])\n        train_sampler = torch.utils.data.RandomSampler(\n            train_dataset,\n            replacement=len(train_idx) < Config.max_samples,\n            num_samples=Config.max_samples\n        )\n        train_loader = DataLoader(train_dataset, \n                                 batch_size=Config.batch_size,\n                                 sampler=train_sampler,\n                                 num_workers=Config.num_workers)\n        \n        val_loader = DataLoader(FOGDataset(X[val_idx], y[val_idx]), \n                              batch_size=Config.batch_size*2,\n                              shuffle=False)\n        \n        # 模型初始化\n        model = ResNet1D().to(Config.device)\n        optimizer = optim.AdamW(model.parameters(), lr=Config.lr)\n        criterion = nn.BCEWithLogitsLoss()\n        scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min')\n        \n        # 训练循环\n        model.train()\n        for epoch in range(Config.num_epochs):\n            for inputs, targets in train_loader:\n                inputs = inputs.to(Config.device)\n                targets = targets.to(Config.device)\n                \n                optimizer.zero_grad()\n                outputs = model(inputs)\n                loss = criterion(outputs, targets)\n                loss.backward()\n                optimizer.step()\n        \n        # 验证和保存\n        model.eval()\n        val_loss = 0.0\n        with torch.no_grad():\n            for inputs, targets in val_loader:\n                outputs = model(inputs.to(Config.device))\n                loss = criterion(outputs, targets.to(Config.device))\n                val_loss += loss.item() * inputs.size(0)\n        \n        val_loss /= len(val_loader.dataset)\n        scheduler.step(val_loss)\n        torch.save(model.state_dict(), f'resnet_fold{fold}.pth')\n        print(f\"Fold {fold+1} Val Loss: {val_loss:.4f}\")\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 主流程\nif __name__ == \"__main__\":\n    X, y, groups = load_and_preprocess()\n    gc.collect()\n    train_fold(X, y, groups)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}