{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":56537,"databundleVersionId":8015876,"sourceType":"competition"}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current sessio","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-05-01T07:04:21.786052Z","iopub.execute_input":"2024-05-01T07:04:21.787093Z","iopub.status.idle":"2024-05-01T07:04:21.803995Z","shell.execute_reply.started":"2024-05-01T07:04:21.787051Z","shell.execute_reply":"2024-05-01T07:04:21.802686Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# csv2parquet","metadata":{}},{"cell_type":"code","source":"! mkdir /kaggle/working/leap-atmospheric-physics-ai-climsim","metadata":{"execution":{"iopub.status.busy":"2024-05-01T07:04:21.80602Z","iopub.execute_input":"2024-05-01T07:04:21.806498Z","iopub.status.idle":"2024-05-01T07:04:22.889618Z","shell.execute_reply.started":"2024-05-01T07:04:21.806462Z","shell.execute_reply":"2024-05-01T07:04:22.888182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n# train.csv: 10,091,520rows, 181.72GB\nimport pandas as pd\nfrom multiprocessing import Pool\nimport os\nfrom tqdm import tqdm\n\nprint('available cpus:', os.cpu_count())\n\n# Define a function to handle processing each chunk\ndef process_chunk(args):\n    chunk_index, chunk = args\n    chunk.to_parquet(f'/kaggle/working/leap-atmospheric-physics-ai-climsim/train_{chunk_index}.parquet')\n\n# Read CSV file using chunksize\ncsv_file_path = '/kaggle/input/leap-atmospheric-physics-ai-climsim/train.csv'\n\n# Calculate the number of chunks based on the total number of rows and chunksize\ntotal_rows = 10_091_520\nchunksize = 100_000\nnum_chunks = (total_rows + chunksize - 1) // chunksize\n\n# Create a list of chunk indices\nchunk_indices = range(num_chunks)\n\n# Use multiprocessing to process chunks in parallel with progress bar\nwith Pool(os.cpu_count()) as pool:\n    with tqdm(total=num_chunks, desc='Processing chunks') as pbar:\n        for _ in pool.imap_unordered(process_chunk, zip(chunk_indices, pd.read_csv(\n            csv_file_path,\n            index_col = 'sample_id',\n            chunksize=chunksize,\n            nrows=chunksize*2, # use only 2 chunksize to check pipeline\n        ))):\n            pbar.update(1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n# test.csv: 625,000rows, 7GB\nimport pandas as pd\npd.read_csv(\n    '/kaggle/input/leap-atmospheric-physics-ai-climsim/test.csv',\n    index_col = 'sample_id',\n).to_parquet('/kaggle/working/leap-atmospheric-physics-ai-climsim/test.parquet')\npd.read_parquet('/kaggle/working/leap-atmospheric-physics-ai-climsim/test.parquet')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n# sample_submission.csv: 625,000rows, 3.85GB\nimport pandas as pd\npd.read_csv(\n    '/kaggle/input/leap-atmospheric-physics-ai-climsim/sample_submission.csv',\n    index_col = 'sample_id',\n).to_parquet('/kaggle/working/leap-atmospheric-physics-ai-climsim/sample_submission.parquet')\npd.read_parquet('/kaggle/working/leap-atmospheric-physics-ai-climsim/sample_submission.parquet')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EDA","metadata":{}},{"cell_type":"markdown","source":"## Input Columns\n\n| Name               | Description                            | Dimension | Units   |\n| ------------------ | -------------------------------------- | --------- | ------- |\n| `state_t`          | air temperature                        | 60        | 𝐾       |\n| `state_q0001`      | specific humidity                      | 60        | 𝑘𝑔/𝑘𝑔   |\n| `state_q0002`      | cloud liquid mixing ratio              | 60        | 𝑘𝑔/𝑘𝑔   |\n| `state_q0003`      | cloud ice mixing ratio                 | 60        | 𝑘𝑔/𝑘𝑔   |\n| `state_u`          | zonal wind speed                       | 60        | 𝑚/𝑠     |\n| `state_v`          | meridional wind speed                  | 60        | 𝑚/𝑠     |\n| `state_ps`         | surface pressure                       | 1         | 𝑃𝑎      |\n| `pbuf_SOLIN`       | solar insolation                       | 1         | 𝑊/𝑚2    |\n| `pbuf_LHFLX`       | surface latent heat flux               | 1         | 𝑊/𝑚2    |\n| `pbuf_SHFLX`       | surface sensible heat flux             | 1         | 𝑊/𝑚2    |\n| `pbuf_TAUX`        | zonal surface stress                   | 1         | 𝑁/𝑚2    |\n| `pbuf_TAUY`        | meridional surface stress              | 1         | 𝑁/𝑚2    |\n| `pbuf_COSZRS`      | cosine of solar zenith angle           | 1         | 𝑁/𝑚2    |\n| `cam_in_ALDIF`     | albedo for diffuse longwave radiation  | 1         |         |\n| `cam_in_ALDIR`     | albedo for direct longwave radiation   | 1         |         |\n| `cam_in_ASDIF`     | albedo for diffuse shortwave radiation | 1         |         |\n| `cam_in_ASDIR`     | albedo for direct shortwave radiation  | 1         |         |\n| `cam_in_LWUP`      | upward longwave flux                   | 1         | 𝑊/𝑚2    |\n| `cam_in_ICEFRAC`   | sea-ice areal fraction                 | 1         |         |\n| `cam_in_LANDFRAC`  | land areal fraction                    | 1         |         |\n| `cam_in_OCNFRAC`   | ocean areal fraction                   | 1         |         |\n| `cam_in_SNOWHLAND` | snow depth over land                   | 1         | 𝑚       |\n| `pbuf_ozone`       | ozone volume mixing ratio              | 60        | 𝑚𝑜𝑙/𝑚𝑜𝑙 |\n| `pbuf_CH4`         | methane volume mixing ratio            | 60        | 𝑚𝑜𝑙/𝑚𝑜𝑙 |\n| `pbuf_N2O`         | nitrous oxide volume mixing ratio      | 60        | 𝑚𝑜𝑙/𝑚𝑜𝑙 |\n\n## Target Columns\n\n| Name             | Description                                          | Dimension | Units   |\n| ---------------- | ---------------------------------------------------- | --------- | ------- |\n| `ptend_t`        | heating tendency                                     | 60        | 𝐾/𝑠     |\n| `ptend_q0001`    | moistening tendency                                  | 60        | 𝑘𝑔/𝑘𝑔/𝑠 |\n| `ptend_q0002`    | cloud liquid mixing ratio change over time           | 60        | 𝑘𝑔/𝑘𝑔/𝑠 |\n| `ptend_q0003`    | cloud ice mixing ratio change over time              | 60        | 𝑘𝑔/𝑘𝑔/𝑠 |\n| `ptend_u`        | zonal wind acceleration                              | 60        | 𝑚/𝑠2    |\n| `ptend_v`        | meridional wind acceleration                         | 60        | 𝑚/𝑠2    |\n| `cam_out_NETSW`  | net shortwave flux at surface                        | 1         | 𝑊/𝑚2    |\n| `cam_out_FLWDS`  | downward longwave flux at surface                    | 1         | 𝑊/𝑚2    |\n| `cam_out_PRECSC` | snow rate (liquid water equivalent)                  | 1         | 𝑚/𝑠     |\n| `cam_out_PRECC`  | rain rate                                            | 1         | 𝑚/𝑠     |\n| `cam_out_SOLS`   | downward visible direct solar flux to surface        | 1         | 𝑊/𝑚2    |\n| `cam_out_SOLL`   | downward near-infrared direct solar flux to surface  | 1         | 𝑊/𝑚2    |\n| `cam_out_SOLSD`  | downward diffuse solar flux to surface               | 1         | 𝑊/𝑚2    |\n| `cam_out_SOLLD`  | downward diffuse near-infrared solar flux to surface | 1         | 𝑊/𝑚2    |","metadata":{}},{"cell_type":"code","source":"index_col = ['sample_id']\n\nfeature_cols = [\n    'state_t_0', 'state_t_1', 'state_t_2', 'state_t_3', 'state_t_4', 'state_t_5', 'state_t_6', 'state_t_7', 'state_t_8', 'state_t_9', 'state_t_10', 'state_t_11', 'state_t_12', 'state_t_13', 'state_t_14', 'state_t_15', 'state_t_16', 'state_t_17', 'state_t_18', 'state_t_19', 'state_t_20', 'state_t_21', 'state_t_22', 'state_t_23', 'state_t_24', 'state_t_25', 'state_t_26', 'state_t_27', 'state_t_28', 'state_t_29', 'state_t_30', 'state_t_31', 'state_t_32', 'state_t_33', 'state_t_34', 'state_t_35', 'state_t_36', 'state_t_37', 'state_t_38', 'state_t_39', 'state_t_40', 'state_t_41', 'state_t_42', 'state_t_43', 'state_t_44', 'state_t_45', 'state_t_46', 'state_t_47', 'state_t_48', 'state_t_49', 'state_t_50', 'state_t_51', 'state_t_52', 'state_t_53', 'state_t_54', 'state_t_55', 'state_t_56', 'state_t_57', 'state_t_58', 'state_t_59',\n    'state_u_0', 'state_u_1', 'state_u_2', 'state_u_3', 'state_u_4', 'state_u_5', 'state_u_6', 'state_u_7', 'state_u_8', 'state_u_9', 'state_u_10', 'state_u_11', 'state_u_12', 'state_u_13', 'state_u_14', 'state_u_15', 'state_u_16', 'state_u_17', 'state_u_18', 'state_u_19', 'state_u_20', 'state_u_21', 'state_u_22', 'state_u_23', 'state_u_24', 'state_u_25', 'state_u_26', 'state_u_27', 'state_u_28', 'state_u_29', 'state_u_30', 'state_u_31', 'state_u_32', 'state_u_33', 'state_u_34', 'state_u_35', 'state_u_36', 'state_u_37', 'state_u_38', 'state_u_39', 'state_u_40', 'state_u_41', 'state_u_42', 'state_u_43', 'state_u_44', 'state_u_45', 'state_u_46', 'state_u_47', 'state_u_48', 'state_u_49', 'state_u_50', 'state_u_51', 'state_u_52', 'state_u_53', 'state_u_54', 'state_u_55', 'state_u_56', 'state_u_57', 'state_u_58', 'state_u_59',\n    'state_v_0', 'state_v_1', 'state_v_2', 'state_v_3', 'state_v_4', 'state_v_5', 'state_v_6', 'state_v_7', 'state_v_8', 'state_v_9', 'state_v_10', 'state_v_11', 'state_v_12', 'state_v_13', 'state_v_14', 'state_v_15', 'state_v_16', 'state_v_17', 'state_v_18', 'state_v_19', 'state_v_20', 'state_v_21', 'state_v_22', 'state_v_23', 'state_v_24', 'state_v_25', 'state_v_26', 'state_v_27', 'state_v_28', 'state_v_29', 'state_v_30', 'state_v_31', 'state_v_32', 'state_v_33', 'state_v_34', 'state_v_35', 'state_v_36', 'state_v_37', 'state_v_38', 'state_v_39', 'state_v_40', 'state_v_41', 'state_v_42', 'state_v_43', 'state_v_44', 'state_v_45', 'state_v_46', 'state_v_47', 'state_v_48', 'state_v_49', 'state_v_50', 'state_v_51', 'state_v_52', 'state_v_53', 'state_v_54', 'state_v_55', 'state_v_56', 'state_v_57', 'state_v_58', 'state_v_59',\n    'state_q0001_0', 'state_q0001_1', 'state_q0001_2', 'state_q0001_3', 'state_q0001_4', 'state_q0001_5', 'state_q0001_6', 'state_q0001_7', 'state_q0001_8', 'state_q0001_9', 'state_q0001_10', 'state_q0001_11', 'state_q0001_12', 'state_q0001_13', 'state_q0001_14', 'state_q0001_15', 'state_q0001_16', 'state_q0001_17', 'state_q0001_18', 'state_q0001_19', 'state_q0001_20', 'state_q0001_21', 'state_q0001_22', 'state_q0001_23', 'state_q0001_24', 'state_q0001_25', 'state_q0001_26', 'state_q0001_27', 'state_q0001_28', 'state_q0001_29', 'state_q0001_30', 'state_q0001_31', 'state_q0001_32', 'state_q0001_33', 'state_q0001_34', 'state_q0001_35', 'state_q0001_36', 'state_q0001_37', 'state_q0001_38', 'state_q0001_39', 'state_q0001_40', 'state_q0001_41', 'state_q0001_42', 'state_q0001_43', 'state_q0001_44', 'state_q0001_45', 'state_q0001_46', 'state_q0001_47', 'state_q0001_48', 'state_q0001_49', 'state_q0001_50', 'state_q0001_51', 'state_q0001_52', 'state_q0001_53', 'state_q0001_54', 'state_q0001_55', 'state_q0001_56', 'state_q0001_57', 'state_q0001_58', 'state_q0001_59',\n    'state_q0002_0', 'state_q0002_1', 'state_q0002_2', 'state_q0002_3', 'state_q0002_4', 'state_q0002_5', 'state_q0002_6', 'state_q0002_7', 'state_q0002_8', 'state_q0002_9', 'state_q0002_10', 'state_q0002_11', 'state_q0002_12', 'state_q0002_13', 'state_q0002_14', 'state_q0002_15', 'state_q0002_16', 'state_q0002_17', 'state_q0002_18', 'state_q0002_19', 'state_q0002_20', 'state_q0002_21', 'state_q0002_22', 'state_q0002_23', 'state_q0002_24', 'state_q0002_25', 'state_q0002_26', 'state_q0002_27', 'state_q0002_28', 'state_q0002_29', 'state_q0002_30', 'state_q0002_31', 'state_q0002_32', 'state_q0002_33', 'state_q0002_34', 'state_q0002_35', 'state_q0002_36', 'state_q0002_37', 'state_q0002_38', 'state_q0002_39', 'state_q0002_40', 'state_q0002_41', 'state_q0002_42', 'state_q0002_43', 'state_q0002_44', 'state_q0002_45', 'state_q0002_46', 'state_q0002_47', 'state_q0002_48', 'state_q0002_49', 'state_q0002_50', 'state_q0002_51', 'state_q0002_52', 'state_q0002_53', 'state_q0002_54', 'state_q0002_55', 'state_q0002_56', 'state_q0002_57', 'state_q0002_58', 'state_q0002_59',\n    'state_q0003_0', 'state_q0003_1', 'state_q0003_2', 'state_q0003_3', 'state_q0003_4', 'state_q0003_5', 'state_q0003_6', 'state_q0003_7', 'state_q0003_8', 'state_q0003_9', 'state_q0003_10', 'state_q0003_11', 'state_q0003_12', 'state_q0003_13', 'state_q0003_14', 'state_q0003_15', 'state_q0003_16', 'state_q0003_17', 'state_q0003_18', 'state_q0003_19', 'state_q0003_20', 'state_q0003_21', 'state_q0003_22', 'state_q0003_23', 'state_q0003_24', 'state_q0003_25', 'state_q0003_26', 'state_q0003_27', 'state_q0003_28', 'state_q0003_29', 'state_q0003_30', 'state_q0003_31', 'state_q0003_32', 'state_q0003_33', 'state_q0003_34', 'state_q0003_35', 'state_q0003_36', 'state_q0003_37', 'state_q0003_38', 'state_q0003_39', 'state_q0003_40', 'state_q0003_41', 'state_q0003_42', 'state_q0003_43', 'state_q0003_44', 'state_q0003_45', 'state_q0003_46', 'state_q0003_47', 'state_q0003_48', 'state_q0003_49', 'state_q0003_50', 'state_q0003_51', 'state_q0003_52', 'state_q0003_53', 'state_q0003_54', 'state_q0003_55', 'state_q0003_56', 'state_q0003_57', 'state_q0003_58', 'state_q0003_59',\n    'state_ps',\n    'cam_in_ALDIF', 'cam_in_ALDIR', 'cam_in_ASDIF', 'cam_in_ASDIR', 'cam_in_LWUP', 'cam_in_ICEFRAC', 'cam_in_LANDFRAC', 'cam_in_OCNFRAC', 'cam_in_SNOWHLAND',\n    'pbuf_SOLIN', 'pbuf_LHFLX', 'pbuf_SHFLX', 'pbuf_TAUX', 'pbuf_TAUY', 'pbuf_COSZRS',\n    'pbuf_ozone_0', 'pbuf_ozone_1', 'pbuf_ozone_2', 'pbuf_ozone_3', 'pbuf_ozone_4', 'pbuf_ozone_5', 'pbuf_ozone_6', 'pbuf_ozone_7', 'pbuf_ozone_8', 'pbuf_ozone_9', 'pbuf_ozone_10', 'pbuf_ozone_11', 'pbuf_ozone_12', 'pbuf_ozone_13', 'pbuf_ozone_14', 'pbuf_ozone_15', 'pbuf_ozone_16', 'pbuf_ozone_17', 'pbuf_ozone_18', 'pbuf_ozone_19', 'pbuf_ozone_20', 'pbuf_ozone_21', 'pbuf_ozone_22', 'pbuf_ozone_23', 'pbuf_ozone_24', 'pbuf_ozone_25', 'pbuf_ozone_26', 'pbuf_ozone_27', 'pbuf_ozone_28', 'pbuf_ozone_29', 'pbuf_ozone_30', 'pbuf_ozone_31', 'pbuf_ozone_32', 'pbuf_ozone_33', 'pbuf_ozone_34', 'pbuf_ozone_35', 'pbuf_ozone_36', 'pbuf_ozone_37', 'pbuf_ozone_38', 'pbuf_ozone_39', 'pbuf_ozone_40', 'pbuf_ozone_41', 'pbuf_ozone_42', 'pbuf_ozone_43', 'pbuf_ozone_44', 'pbuf_ozone_45', 'pbuf_ozone_46', 'pbuf_ozone_47', 'pbuf_ozone_48', 'pbuf_ozone_49', 'pbuf_ozone_50', 'pbuf_ozone_51', 'pbuf_ozone_52', 'pbuf_ozone_53', 'pbuf_ozone_54', 'pbuf_ozone_55', 'pbuf_ozone_56', 'pbuf_ozone_57', 'pbuf_ozone_58', 'pbuf_ozone_59',\n    'pbuf_CH4_0', 'pbuf_CH4_1', 'pbuf_CH4_2', 'pbuf_CH4_3', 'pbuf_CH4_4', 'pbuf_CH4_5', 'pbuf_CH4_6', 'pbuf_CH4_7', 'pbuf_CH4_8', 'pbuf_CH4_9', 'pbuf_CH4_10', 'pbuf_CH4_11', 'pbuf_CH4_12', 'pbuf_CH4_13', 'pbuf_CH4_14', 'pbuf_CH4_15', 'pbuf_CH4_16', 'pbuf_CH4_17', 'pbuf_CH4_18', 'pbuf_CH4_19', 'pbuf_CH4_20', 'pbuf_CH4_21', 'pbuf_CH4_22', 'pbuf_CH4_23', 'pbuf_CH4_24', 'pbuf_CH4_25', 'pbuf_CH4_26', 'pbuf_CH4_27', 'pbuf_CH4_28', 'pbuf_CH4_29', 'pbuf_CH4_30', 'pbuf_CH4_31', 'pbuf_CH4_32', 'pbuf_CH4_33', 'pbuf_CH4_34', 'pbuf_CH4_35', 'pbuf_CH4_36', 'pbuf_CH4_37', 'pbuf_CH4_38', 'pbuf_CH4_39', 'pbuf_CH4_40', 'pbuf_CH4_41', 'pbuf_CH4_42', 'pbuf_CH4_43', 'pbuf_CH4_44', 'pbuf_CH4_45', 'pbuf_CH4_46', 'pbuf_CH4_47', 'pbuf_CH4_48', 'pbuf_CH4_49', 'pbuf_CH4_50', 'pbuf_CH4_51', 'pbuf_CH4_52', 'pbuf_CH4_53', 'pbuf_CH4_54', 'pbuf_CH4_55', 'pbuf_CH4_56', 'pbuf_CH4_57', 'pbuf_CH4_58', 'pbuf_CH4_59',\n    'pbuf_N2O_0', 'pbuf_N2O_1', 'pbuf_N2O_2', 'pbuf_N2O_3', 'pbuf_N2O_4', 'pbuf_N2O_5', 'pbuf_N2O_6', 'pbuf_N2O_7', 'pbuf_N2O_8', 'pbuf_N2O_9', 'pbuf_N2O_10', 'pbuf_N2O_11', 'pbuf_N2O_12', 'pbuf_N2O_13', 'pbuf_N2O_14', 'pbuf_N2O_15', 'pbuf_N2O_16', 'pbuf_N2O_17', 'pbuf_N2O_18', 'pbuf_N2O_19', 'pbuf_N2O_20', 'pbuf_N2O_21', 'pbuf_N2O_22', 'pbuf_N2O_23', 'pbuf_N2O_24', 'pbuf_N2O_25', 'pbuf_N2O_26', 'pbuf_N2O_27', 'pbuf_N2O_28', 'pbuf_N2O_29', 'pbuf_N2O_30', 'pbuf_N2O_31', 'pbuf_N2O_32', 'pbuf_N2O_33', 'pbuf_N2O_34', 'pbuf_N2O_35', 'pbuf_N2O_36', 'pbuf_N2O_37', 'pbuf_N2O_38', 'pbuf_N2O_39', 'pbuf_N2O_40', 'pbuf_N2O_41', 'pbuf_N2O_42', 'pbuf_N2O_43', 'pbuf_N2O_44', 'pbuf_N2O_45', 'pbuf_N2O_46', 'pbuf_N2O_47', 'pbuf_N2O_48', 'pbuf_N2O_49', 'pbuf_N2O_50', 'pbuf_N2O_51', 'pbuf_N2O_52', 'pbuf_N2O_53', 'pbuf_N2O_54', 'pbuf_N2O_55', 'pbuf_N2O_56', 'pbuf_N2O_57', 'pbuf_N2O_58', 'pbuf_N2O_59'\n]\nprint('feature_cols:', len(feature_cols))\n\ntarget_cols = [\n    'ptend_t_0', 'ptend_t_1', 'ptend_t_2', 'ptend_t_3', 'ptend_t_4', 'ptend_t_5', 'ptend_t_6', 'ptend_t_7', 'ptend_t_8', 'ptend_t_9', 'ptend_t_10', 'ptend_t_11', 'ptend_t_12', 'ptend_t_13', 'ptend_t_14', 'ptend_t_15', 'ptend_t_16', 'ptend_t_17', 'ptend_t_18', 'ptend_t_19', 'ptend_t_20', 'ptend_t_21', 'ptend_t_22', 'ptend_t_23', 'ptend_t_24', 'ptend_t_25', 'ptend_t_26', 'ptend_t_27', 'ptend_t_28', 'ptend_t_29', 'ptend_t_30', 'ptend_t_31', 'ptend_t_32', 'ptend_t_33', 'ptend_t_34', 'ptend_t_35', 'ptend_t_36', 'ptend_t_37', 'ptend_t_38', 'ptend_t_39', 'ptend_t_40', 'ptend_t_41', 'ptend_t_42', 'ptend_t_43', 'ptend_t_44', 'ptend_t_45', 'ptend_t_46', 'ptend_t_47', 'ptend_t_48', 'ptend_t_49', 'ptend_t_50', 'ptend_t_51', 'ptend_t_52', 'ptend_t_53', 'ptend_t_54', 'ptend_t_55', 'ptend_t_56', 'ptend_t_57', 'ptend_t_58', 'ptend_t_59',\n    'ptend_u_0', 'ptend_u_1', 'ptend_u_2', 'ptend_u_3', 'ptend_u_4', 'ptend_u_5', 'ptend_u_6', 'ptend_u_7', 'ptend_u_8', 'ptend_u_9', 'ptend_u_10', 'ptend_u_11', 'ptend_u_12', 'ptend_u_13', 'ptend_u_14', 'ptend_u_15', 'ptend_u_16', 'ptend_u_17', 'ptend_u_18', 'ptend_u_19', 'ptend_u_20', 'ptend_u_21', 'ptend_u_22', 'ptend_u_23', 'ptend_u_24', 'ptend_u_25', 'ptend_u_26', 'ptend_u_27', 'ptend_u_28', 'ptend_u_29', 'ptend_u_30', 'ptend_u_31', 'ptend_u_32', 'ptend_u_33', 'ptend_u_34', 'ptend_u_35', 'ptend_u_36', 'ptend_u_37', 'ptend_u_38', 'ptend_u_39', 'ptend_u_40', 'ptend_u_41', 'ptend_u_42', 'ptend_u_43', 'ptend_u_44', 'ptend_u_45', 'ptend_u_46', 'ptend_u_47', 'ptend_u_48', 'ptend_u_49', 'ptend_u_50', 'ptend_u_51', 'ptend_u_52', 'ptend_u_53', 'ptend_u_54', 'ptend_u_55', 'ptend_u_56', 'ptend_u_57', 'ptend_u_58', 'ptend_u_59',\n    'ptend_v_0', 'ptend_v_1', 'ptend_v_2', 'ptend_v_3', 'ptend_v_4', 'ptend_v_5', 'ptend_v_6', 'ptend_v_7', 'ptend_v_8', 'ptend_v_9', 'ptend_v_10', 'ptend_v_11', 'ptend_v_12', 'ptend_v_13', 'ptend_v_14', 'ptend_v_15', 'ptend_v_16', 'ptend_v_17', 'ptend_v_18', 'ptend_v_19', 'ptend_v_20', 'ptend_v_21', 'ptend_v_22', 'ptend_v_23', 'ptend_v_24', 'ptend_v_25', 'ptend_v_26', 'ptend_v_27', 'ptend_v_28', 'ptend_v_29', 'ptend_v_30', 'ptend_v_31', 'ptend_v_32', 'ptend_v_33', 'ptend_v_34', 'ptend_v_35', 'ptend_v_36', 'ptend_v_37', 'ptend_v_38', 'ptend_v_39', 'ptend_v_40', 'ptend_v_41', 'ptend_v_42', 'ptend_v_43', 'ptend_v_44', 'ptend_v_45', 'ptend_v_46', 'ptend_v_47', 'ptend_v_48', 'ptend_v_49', 'ptend_v_50', 'ptend_v_51', 'ptend_v_52', 'ptend_v_53', 'ptend_v_54', 'ptend_v_55', 'ptend_v_56', 'ptend_v_57', 'ptend_v_58', 'ptend_v_59',\n    'ptend_q0001_0', 'ptend_q0001_1', 'ptend_q0001_2', 'ptend_q0001_3', 'ptend_q0001_4', 'ptend_q0001_5', 'ptend_q0001_6', 'ptend_q0001_7', 'ptend_q0001_8', 'ptend_q0001_9', 'ptend_q0001_10', 'ptend_q0001_11', 'ptend_q0001_12', 'ptend_q0001_13', 'ptend_q0001_14', 'ptend_q0001_15', 'ptend_q0001_16', 'ptend_q0001_17', 'ptend_q0001_18', 'ptend_q0001_19', 'ptend_q0001_20', 'ptend_q0001_21', 'ptend_q0001_22', 'ptend_q0001_23', 'ptend_q0001_24', 'ptend_q0001_25', 'ptend_q0001_26', 'ptend_q0001_27', 'ptend_q0001_28', 'ptend_q0001_29', 'ptend_q0001_30', 'ptend_q0001_31', 'ptend_q0001_32', 'ptend_q0001_33', 'ptend_q0001_34', 'ptend_q0001_35', 'ptend_q0001_36', 'ptend_q0001_37', 'ptend_q0001_38', 'ptend_q0001_39', 'ptend_q0001_40', 'ptend_q0001_41', 'ptend_q0001_42', 'ptend_q0001_43', 'ptend_q0001_44', 'ptend_q0001_45', 'ptend_q0001_46', 'ptend_q0001_47', 'ptend_q0001_48', 'ptend_q0001_49', 'ptend_q0001_50', 'ptend_q0001_51', 'ptend_q0001_52', 'ptend_q0001_53', 'ptend_q0001_54', 'ptend_q0001_55', 'ptend_q0001_56', 'ptend_q0001_57', 'ptend_q0001_58', 'ptend_q0001_59',\n    'ptend_q0002_0', 'ptend_q0002_1', 'ptend_q0002_2', 'ptend_q0002_3', 'ptend_q0002_4', 'ptend_q0002_5', 'ptend_q0002_6', 'ptend_q0002_7', 'ptend_q0002_8', 'ptend_q0002_9', 'ptend_q0002_10', 'ptend_q0002_11', 'ptend_q0002_12', 'ptend_q0002_13', 'ptend_q0002_14', 'ptend_q0002_15', 'ptend_q0002_16', 'ptend_q0002_17', 'ptend_q0002_18', 'ptend_q0002_19', 'ptend_q0002_20', 'ptend_q0002_21', 'ptend_q0002_22', 'ptend_q0002_23', 'ptend_q0002_24', 'ptend_q0002_25', 'ptend_q0002_26', 'ptend_q0002_27', 'ptend_q0002_28', 'ptend_q0002_29', 'ptend_q0002_30', 'ptend_q0002_31', 'ptend_q0002_32', 'ptend_q0002_33', 'ptend_q0002_34', 'ptend_q0002_35', 'ptend_q0002_36', 'ptend_q0002_37', 'ptend_q0002_38', 'ptend_q0002_39', 'ptend_q0002_40', 'ptend_q0002_41', 'ptend_q0002_42', 'ptend_q0002_43', 'ptend_q0002_44', 'ptend_q0002_45', 'ptend_q0002_46', 'ptend_q0002_47', 'ptend_q0002_48', 'ptend_q0002_49', 'ptend_q0002_50', 'ptend_q0002_51', 'ptend_q0002_52', 'ptend_q0002_53', 'ptend_q0002_54', 'ptend_q0002_55', 'ptend_q0002_56', 'ptend_q0002_57', 'ptend_q0002_58', 'ptend_q0002_59',\n    'ptend_q0003_0', 'ptend_q0003_1', 'ptend_q0003_2', 'ptend_q0003_3', 'ptend_q0003_4', 'ptend_q0003_5', 'ptend_q0003_6', 'ptend_q0003_7', 'ptend_q0003_8', 'ptend_q0003_9', 'ptend_q0003_10', 'ptend_q0003_11', 'ptend_q0003_12', 'ptend_q0003_13', 'ptend_q0003_14', 'ptend_q0003_15', 'ptend_q0003_16', 'ptend_q0003_17', 'ptend_q0003_18', 'ptend_q0003_19', 'ptend_q0003_20', 'ptend_q0003_21', 'ptend_q0003_22', 'ptend_q0003_23', 'ptend_q0003_24', 'ptend_q0003_25', 'ptend_q0003_26', 'ptend_q0003_27', 'ptend_q0003_28', 'ptend_q0003_29', 'ptend_q0003_30', 'ptend_q0003_31', 'ptend_q0003_32', 'ptend_q0003_33', 'ptend_q0003_34', 'ptend_q0003_35', 'ptend_q0003_36', 'ptend_q0003_37', 'ptend_q0003_38', 'ptend_q0003_39', 'ptend_q0003_40', 'ptend_q0003_41', 'ptend_q0003_42', 'ptend_q0003_43', 'ptend_q0003_44', 'ptend_q0003_45', 'ptend_q0003_46', 'ptend_q0003_47', 'ptend_q0003_48', 'ptend_q0003_49', 'ptend_q0003_50', 'ptend_q0003_51', 'ptend_q0003_52', 'ptend_q0003_53', 'ptend_q0003_54', 'ptend_q0003_55', 'ptend_q0003_56', 'ptend_q0003_57', 'ptend_q0003_58', 'ptend_q0003_59',\n    'cam_out_NETSW', 'cam_out_FLWDS', 'cam_out_PRECSC', 'cam_out_PRECC', 'cam_out_SOLS', 'cam_out_SOLL', 'cam_out_SOLSD', 'cam_out_SOLLD'\n]\nprint('target_cols:', len(target_cols))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"feature_groups = {\n    'state_t': ['state_t_0', 'state_t_1', 'state_t_2', 'state_t_3', 'state_t_4', 'state_t_5', 'state_t_6', 'state_t_7', 'state_t_8', 'state_t_9', 'state_t_10', 'state_t_11', 'state_t_12', 'state_t_13', 'state_t_14', 'state_t_15', 'state_t_16', 'state_t_17', 'state_t_18', 'state_t_19', 'state_t_20', 'state_t_21', 'state_t_22', 'state_t_23', 'state_t_24', 'state_t_25', 'state_t_26', 'state_t_27', 'state_t_28', 'state_t_29', 'state_t_30', 'state_t_31', 'state_t_32', 'state_t_33', 'state_t_34', 'state_t_35', 'state_t_36', 'state_t_37', 'state_t_38', 'state_t_39', 'state_t_40', 'state_t_41', 'state_t_42', 'state_t_43', 'state_t_44', 'state_t_45', 'state_t_46', 'state_t_47', 'state_t_48', 'state_t_49', 'state_t_50', 'state_t_51', 'state_t_52', 'state_t_53', 'state_t_54', 'state_t_55', 'state_t_56', 'state_t_57', 'state_t_58', 'state_t_59',],\n    'state_u': ['state_u_0', 'state_u_1', 'state_u_2', 'state_u_3', 'state_u_4', 'state_u_5', 'state_u_6', 'state_u_7', 'state_u_8', 'state_u_9', 'state_u_10', 'state_u_11', 'state_u_12', 'state_u_13', 'state_u_14', 'state_u_15', 'state_u_16', 'state_u_17', 'state_u_18', 'state_u_19', 'state_u_20', 'state_u_21', 'state_u_22', 'state_u_23', 'state_u_24', 'state_u_25', 'state_u_26', 'state_u_27', 'state_u_28', 'state_u_29', 'state_u_30', 'state_u_31', 'state_u_32', 'state_u_33', 'state_u_34', 'state_u_35', 'state_u_36', 'state_u_37', 'state_u_38', 'state_u_39', 'state_u_40', 'state_u_41', 'state_u_42', 'state_u_43', 'state_u_44', 'state_u_45', 'state_u_46', 'state_u_47', 'state_u_48', 'state_u_49', 'state_u_50', 'state_u_51', 'state_u_52', 'state_u_53', 'state_u_54', 'state_u_55', 'state_u_56', 'state_u_57', 'state_u_58', 'state_u_59',],\n    'state_v': ['state_v_0', 'state_v_1', 'state_v_2', 'state_v_3', 'state_v_4', 'state_v_5', 'state_v_6', 'state_v_7', 'state_v_8', 'state_v_9', 'state_v_10', 'state_v_11', 'state_v_12', 'state_v_13', 'state_v_14', 'state_v_15', 'state_v_16', 'state_v_17', 'state_v_18', 'state_v_19', 'state_v_20', 'state_v_21', 'state_v_22', 'state_v_23', 'state_v_24', 'state_v_25', 'state_v_26', 'state_v_27', 'state_v_28', 'state_v_29', 'state_v_30', 'state_v_31', 'state_v_32', 'state_v_33', 'state_v_34', 'state_v_35', 'state_v_36', 'state_v_37', 'state_v_38', 'state_v_39', 'state_v_40', 'state_v_41', 'state_v_42', 'state_v_43', 'state_v_44', 'state_v_45', 'state_v_46', 'state_v_47', 'state_v_48', 'state_v_49', 'state_v_50', 'state_v_51', 'state_v_52', 'state_v_53', 'state_v_54', 'state_v_55', 'state_v_56', 'state_v_57', 'state_v_58', 'state_v_59',],\n    'state_q0001': ['state_q0001_0', 'state_q0001_1', 'state_q0001_2', 'state_q0001_3', 'state_q0001_4', 'state_q0001_5', 'state_q0001_6', 'state_q0001_7', 'state_q0001_8', 'state_q0001_9', 'state_q0001_10', 'state_q0001_11', 'state_q0001_12', 'state_q0001_13', 'state_q0001_14', 'state_q0001_15', 'state_q0001_16', 'state_q0001_17', 'state_q0001_18', 'state_q0001_19', 'state_q0001_20', 'state_q0001_21', 'state_q0001_22', 'state_q0001_23', 'state_q0001_24', 'state_q0001_25', 'state_q0001_26', 'state_q0001_27', 'state_q0001_28', 'state_q0001_29', 'state_q0001_30', 'state_q0001_31', 'state_q0001_32', 'state_q0001_33', 'state_q0001_34', 'state_q0001_35', 'state_q0001_36', 'state_q0001_37', 'state_q0001_38', 'state_q0001_39', 'state_q0001_40', 'state_q0001_41', 'state_q0001_42', 'state_q0001_43', 'state_q0001_44', 'state_q0001_45', 'state_q0001_46', 'state_q0001_47', 'state_q0001_48', 'state_q0001_49', 'state_q0001_50', 'state_q0001_51', 'state_q0001_52', 'state_q0001_53', 'state_q0001_54', 'state_q0001_55', 'state_q0001_56', 'state_q0001_57', 'state_q0001_58', 'state_q0001_59',],\n    'state_q0002': ['state_q0002_0', 'state_q0002_1', 'state_q0002_2', 'state_q0002_3', 'state_q0002_4', 'state_q0002_5', 'state_q0002_6', 'state_q0002_7', 'state_q0002_8', 'state_q0002_9', 'state_q0002_10', 'state_q0002_11', 'state_q0002_12', 'state_q0002_13', 'state_q0002_14', 'state_q0002_15', 'state_q0002_16', 'state_q0002_17', 'state_q0002_18', 'state_q0002_19', 'state_q0002_20', 'state_q0002_21', 'state_q0002_22', 'state_q0002_23', 'state_q0002_24', 'state_q0002_25', 'state_q0002_26', 'state_q0002_27', 'state_q0002_28', 'state_q0002_29', 'state_q0002_30', 'state_q0002_31', 'state_q0002_32', 'state_q0002_33', 'state_q0002_34', 'state_q0002_35', 'state_q0002_36', 'state_q0002_37', 'state_q0002_38', 'state_q0002_39', 'state_q0002_40', 'state_q0002_41', 'state_q0002_42', 'state_q0002_43', 'state_q0002_44', 'state_q0002_45', 'state_q0002_46', 'state_q0002_47', 'state_q0002_48', 'state_q0002_49', 'state_q0002_50', 'state_q0002_51', 'state_q0002_52', 'state_q0002_53', 'state_q0002_54', 'state_q0002_55', 'state_q0002_56', 'state_q0002_57', 'state_q0002_58', 'state_q0002_59',],\n    'state_q0003': ['state_q0003_0', 'state_q0003_1', 'state_q0003_2', 'state_q0003_3', 'state_q0003_4', 'state_q0003_5', 'state_q0003_6', 'state_q0003_7', 'state_q0003_8', 'state_q0003_9', 'state_q0003_10', 'state_q0003_11', 'state_q0003_12', 'state_q0003_13', 'state_q0003_14', 'state_q0003_15', 'state_q0003_16', 'state_q0003_17', 'state_q0003_18', 'state_q0003_19', 'state_q0003_20', 'state_q0003_21', 'state_q0003_22', 'state_q0003_23', 'state_q0003_24', 'state_q0003_25', 'state_q0003_26', 'state_q0003_27', 'state_q0003_28', 'state_q0003_29', 'state_q0003_30', 'state_q0003_31', 'state_q0003_32', 'state_q0003_33', 'state_q0003_34', 'state_q0003_35', 'state_q0003_36', 'state_q0003_37', 'state_q0003_38', 'state_q0003_39', 'state_q0003_40', 'state_q0003_41', 'state_q0003_42', 'state_q0003_43', 'state_q0003_44', 'state_q0003_45', 'state_q0003_46', 'state_q0003_47', 'state_q0003_48', 'state_q0003_49', 'state_q0003_50', 'state_q0003_51', 'state_q0003_52', 'state_q0003_53', 'state_q0003_54', 'state_q0003_55', 'state_q0003_56', 'state_q0003_57', 'state_q0003_58', 'state_q0003_59',],\n    'state_ps': ['state_ps',],\n    'cam': ['cam_in_ALDIF', 'cam_in_ALDIR', 'cam_in_ASDIF', 'cam_in_ASDIR', 'cam_in_LWUP', 'cam_in_ICEFRAC', 'cam_in_LANDFRAC', 'cam_in_OCNFRAC', 'cam_in_SNOWHLAND',],\n    'pbuf': ['pbuf_SOLIN', 'pbuf_LHFLX', 'pbuf_SHFLX', 'pbuf_TAUX', 'pbuf_TAUY', 'pbuf_COSZRS',],\n    'pbuf_ozone': ['pbuf_ozone_0', 'pbuf_ozone_1', 'pbuf_ozone_2', 'pbuf_ozone_3', 'pbuf_ozone_4', 'pbuf_ozone_5', 'pbuf_ozone_6', 'pbuf_ozone_7', 'pbuf_ozone_8', 'pbuf_ozone_9', 'pbuf_ozone_10', 'pbuf_ozone_11', 'pbuf_ozone_12', 'pbuf_ozone_13', 'pbuf_ozone_14', 'pbuf_ozone_15', 'pbuf_ozone_16', 'pbuf_ozone_17', 'pbuf_ozone_18', 'pbuf_ozone_19', 'pbuf_ozone_20', 'pbuf_ozone_21', 'pbuf_ozone_22', 'pbuf_ozone_23', 'pbuf_ozone_24', 'pbuf_ozone_25', 'pbuf_ozone_26', 'pbuf_ozone_27', 'pbuf_ozone_28', 'pbuf_ozone_29', 'pbuf_ozone_30', 'pbuf_ozone_31', 'pbuf_ozone_32', 'pbuf_ozone_33', 'pbuf_ozone_34', 'pbuf_ozone_35', 'pbuf_ozone_36', 'pbuf_ozone_37', 'pbuf_ozone_38', 'pbuf_ozone_39', 'pbuf_ozone_40', 'pbuf_ozone_41', 'pbuf_ozone_42', 'pbuf_ozone_43', 'pbuf_ozone_44', 'pbuf_ozone_45', 'pbuf_ozone_46', 'pbuf_ozone_47', 'pbuf_ozone_48', 'pbuf_ozone_49', 'pbuf_ozone_50', 'pbuf_ozone_51', 'pbuf_ozone_52', 'pbuf_ozone_53', 'pbuf_ozone_54', 'pbuf_ozone_55', 'pbuf_ozone_56', 'pbuf_ozone_57', 'pbuf_ozone_58', 'pbuf_ozone_59',],\n    'pbuf_CH4': ['pbuf_CH4_0', 'pbuf_CH4_1', 'pbuf_CH4_2', 'pbuf_CH4_3', 'pbuf_CH4_4', 'pbuf_CH4_5', 'pbuf_CH4_6', 'pbuf_CH4_7', 'pbuf_CH4_8', 'pbuf_CH4_9', 'pbuf_CH4_10', 'pbuf_CH4_11', 'pbuf_CH4_12', 'pbuf_CH4_13', 'pbuf_CH4_14', 'pbuf_CH4_15', 'pbuf_CH4_16', 'pbuf_CH4_17', 'pbuf_CH4_18', 'pbuf_CH4_19', 'pbuf_CH4_20', 'pbuf_CH4_21', 'pbuf_CH4_22', 'pbuf_CH4_23', 'pbuf_CH4_24', 'pbuf_CH4_25', 'pbuf_CH4_26', 'pbuf_CH4_27', 'pbuf_CH4_28', 'pbuf_CH4_29', 'pbuf_CH4_30', 'pbuf_CH4_31', 'pbuf_CH4_32', 'pbuf_CH4_33', 'pbuf_CH4_34', 'pbuf_CH4_35', 'pbuf_CH4_36', 'pbuf_CH4_37', 'pbuf_CH4_38', 'pbuf_CH4_39', 'pbuf_CH4_40', 'pbuf_CH4_41', 'pbuf_CH4_42', 'pbuf_CH4_43', 'pbuf_CH4_44', 'pbuf_CH4_45', 'pbuf_CH4_46', 'pbuf_CH4_47', 'pbuf_CH4_48', 'pbuf_CH4_49', 'pbuf_CH4_50', 'pbuf_CH4_51', 'pbuf_CH4_52', 'pbuf_CH4_53', 'pbuf_CH4_54', 'pbuf_CH4_55', 'pbuf_CH4_56', 'pbuf_CH4_57', 'pbuf_CH4_58', 'pbuf_CH4_59',],\n    'pbuf_N2O': ['pbuf_N2O_0', 'pbuf_N2O_1', 'pbuf_N2O_2', 'pbuf_N2O_3', 'pbuf_N2O_4', 'pbuf_N2O_5', 'pbuf_N2O_6', 'pbuf_N2O_7', 'pbuf_N2O_8', 'pbuf_N2O_9', 'pbuf_N2O_10', 'pbuf_N2O_11', 'pbuf_N2O_12', 'pbuf_N2O_13', 'pbuf_N2O_14', 'pbuf_N2O_15', 'pbuf_N2O_16', 'pbuf_N2O_17', 'pbuf_N2O_18', 'pbuf_N2O_19', 'pbuf_N2O_20', 'pbuf_N2O_21', 'pbuf_N2O_22', 'pbuf_N2O_23', 'pbuf_N2O_24', 'pbuf_N2O_25', 'pbuf_N2O_26', 'pbuf_N2O_27', 'pbuf_N2O_28', 'pbuf_N2O_29', 'pbuf_N2O_30', 'pbuf_N2O_31', 'pbuf_N2O_32', 'pbuf_N2O_33', 'pbuf_N2O_34', 'pbuf_N2O_35', 'pbuf_N2O_36', 'pbuf_N2O_37', 'pbuf_N2O_38', 'pbuf_N2O_39', 'pbuf_N2O_40', 'pbuf_N2O_41', 'pbuf_N2O_42', 'pbuf_N2O_43', 'pbuf_N2O_44', 'pbuf_N2O_45', 'pbuf_N2O_46', 'pbuf_N2O_47', 'pbuf_N2O_48', 'pbuf_N2O_49', 'pbuf_N2O_50', 'pbuf_N2O_51', 'pbuf_N2O_52', 'pbuf_N2O_53', 'pbuf_N2O_54', 'pbuf_N2O_55', 'pbuf_N2O_56', 'pbuf_N2O_57', 'pbuf_N2O_58', 'pbuf_N2O_59'],\n}\nprint('feature_groups:', feature_groups.keys())\n\ntarget_groups = {\n    'ptend_t': ['ptend_t_0', 'ptend_t_1', 'ptend_t_2', 'ptend_t_3', 'ptend_t_4', 'ptend_t_5', 'ptend_t_6', 'ptend_t_7', 'ptend_t_8', 'ptend_t_9', 'ptend_t_10', 'ptend_t_11', 'ptend_t_12', 'ptend_t_13', 'ptend_t_14', 'ptend_t_15', 'ptend_t_16', 'ptend_t_17', 'ptend_t_18', 'ptend_t_19', 'ptend_t_20', 'ptend_t_21', 'ptend_t_22', 'ptend_t_23', 'ptend_t_24', 'ptend_t_25', 'ptend_t_26', 'ptend_t_27', 'ptend_t_28', 'ptend_t_29', 'ptend_t_30', 'ptend_t_31', 'ptend_t_32', 'ptend_t_33', 'ptend_t_34', 'ptend_t_35', 'ptend_t_36', 'ptend_t_37', 'ptend_t_38', 'ptend_t_39', 'ptend_t_40', 'ptend_t_41', 'ptend_t_42', 'ptend_t_43', 'ptend_t_44', 'ptend_t_45', 'ptend_t_46', 'ptend_t_47', 'ptend_t_48', 'ptend_t_49', 'ptend_t_50', 'ptend_t_51', 'ptend_t_52', 'ptend_t_53', 'ptend_t_54', 'ptend_t_55', 'ptend_t_56', 'ptend_t_57', 'ptend_t_58', 'ptend_t_59',],\n    'ptend_u': ['ptend_u_0', 'ptend_u_1', 'ptend_u_2', 'ptend_u_3', 'ptend_u_4', 'ptend_u_5', 'ptend_u_6', 'ptend_u_7', 'ptend_u_8', 'ptend_u_9', 'ptend_u_10', 'ptend_u_11', 'ptend_u_12', 'ptend_u_13', 'ptend_u_14', 'ptend_u_15', 'ptend_u_16', 'ptend_u_17', 'ptend_u_18', 'ptend_u_19', 'ptend_u_20', 'ptend_u_21', 'ptend_u_22', 'ptend_u_23', 'ptend_u_24', 'ptend_u_25', 'ptend_u_26', 'ptend_u_27', 'ptend_u_28', 'ptend_u_29', 'ptend_u_30', 'ptend_u_31', 'ptend_u_32', 'ptend_u_33', 'ptend_u_34', 'ptend_u_35', 'ptend_u_36', 'ptend_u_37', 'ptend_u_38', 'ptend_u_39', 'ptend_u_40', 'ptend_u_41', 'ptend_u_42', 'ptend_u_43', 'ptend_u_44', 'ptend_u_45', 'ptend_u_46', 'ptend_u_47', 'ptend_u_48', 'ptend_u_49', 'ptend_u_50', 'ptend_u_51', 'ptend_u_52', 'ptend_u_53', 'ptend_u_54', 'ptend_u_55', 'ptend_u_56', 'ptend_u_57', 'ptend_u_58', 'ptend_u_59',],\n    'ptend_v': ['ptend_v_0', 'ptend_v_1', 'ptend_v_2', 'ptend_v_3', 'ptend_v_4', 'ptend_v_5', 'ptend_v_6', 'ptend_v_7', 'ptend_v_8', 'ptend_v_9', 'ptend_v_10', 'ptend_v_11', 'ptend_v_12', 'ptend_v_13', 'ptend_v_14', 'ptend_v_15', 'ptend_v_16', 'ptend_v_17', 'ptend_v_18', 'ptend_v_19', 'ptend_v_20', 'ptend_v_21', 'ptend_v_22', 'ptend_v_23', 'ptend_v_24', 'ptend_v_25', 'ptend_v_26', 'ptend_v_27', 'ptend_v_28', 'ptend_v_29', 'ptend_v_30', 'ptend_v_31', 'ptend_v_32', 'ptend_v_33', 'ptend_v_34', 'ptend_v_35', 'ptend_v_36', 'ptend_v_37', 'ptend_v_38', 'ptend_v_39', 'ptend_v_40', 'ptend_v_41', 'ptend_v_42', 'ptend_v_43', 'ptend_v_44', 'ptend_v_45', 'ptend_v_46', 'ptend_v_47', 'ptend_v_48', 'ptend_v_49', 'ptend_v_50', 'ptend_v_51', 'ptend_v_52', 'ptend_v_53', 'ptend_v_54', 'ptend_v_55', 'ptend_v_56', 'ptend_v_57', 'ptend_v_58', 'ptend_v_59',],\n    'ptend_q0001': ['ptend_q0001_0', 'ptend_q0001_1', 'ptend_q0001_2', 'ptend_q0001_3', 'ptend_q0001_4', 'ptend_q0001_5', 'ptend_q0001_6', 'ptend_q0001_7', 'ptend_q0001_8', 'ptend_q0001_9', 'ptend_q0001_10', 'ptend_q0001_11', 'ptend_q0001_12', 'ptend_q0001_13', 'ptend_q0001_14', 'ptend_q0001_15', 'ptend_q0001_16', 'ptend_q0001_17', 'ptend_q0001_18', 'ptend_q0001_19', 'ptend_q0001_20', 'ptend_q0001_21', 'ptend_q0001_22', 'ptend_q0001_23', 'ptend_q0001_24', 'ptend_q0001_25', 'ptend_q0001_26', 'ptend_q0001_27', 'ptend_q0001_28', 'ptend_q0001_29', 'ptend_q0001_30', 'ptend_q0001_31', 'ptend_q0001_32', 'ptend_q0001_33', 'ptend_q0001_34', 'ptend_q0001_35', 'ptend_q0001_36', 'ptend_q0001_37', 'ptend_q0001_38', 'ptend_q0001_39', 'ptend_q0001_40', 'ptend_q0001_41', 'ptend_q0001_42', 'ptend_q0001_43', 'ptend_q0001_44', 'ptend_q0001_45', 'ptend_q0001_46', 'ptend_q0001_47', 'ptend_q0001_48', 'ptend_q0001_49', 'ptend_q0001_50', 'ptend_q0001_51', 'ptend_q0001_52', 'ptend_q0001_53', 'ptend_q0001_54', 'ptend_q0001_55', 'ptend_q0001_56', 'ptend_q0001_57', 'ptend_q0001_58', 'ptend_q0001_59',],\n    'ptend_q0002': ['ptend_q0002_0', 'ptend_q0002_1', 'ptend_q0002_2', 'ptend_q0002_3', 'ptend_q0002_4', 'ptend_q0002_5', 'ptend_q0002_6', 'ptend_q0002_7', 'ptend_q0002_8', 'ptend_q0002_9', 'ptend_q0002_10', 'ptend_q0002_11', 'ptend_q0002_12', 'ptend_q0002_13', 'ptend_q0002_14', 'ptend_q0002_15', 'ptend_q0002_16', 'ptend_q0002_17', 'ptend_q0002_18', 'ptend_q0002_19', 'ptend_q0002_20', 'ptend_q0002_21', 'ptend_q0002_22', 'ptend_q0002_23', 'ptend_q0002_24', 'ptend_q0002_25', 'ptend_q0002_26', 'ptend_q0002_27', 'ptend_q0002_28', 'ptend_q0002_29', 'ptend_q0002_30', 'ptend_q0002_31', 'ptend_q0002_32', 'ptend_q0002_33', 'ptend_q0002_34', 'ptend_q0002_35', 'ptend_q0002_36', 'ptend_q0002_37', 'ptend_q0002_38', 'ptend_q0002_39', 'ptend_q0002_40', 'ptend_q0002_41', 'ptend_q0002_42', 'ptend_q0002_43', 'ptend_q0002_44', 'ptend_q0002_45', 'ptend_q0002_46', 'ptend_q0002_47', 'ptend_q0002_48', 'ptend_q0002_49', 'ptend_q0002_50', 'ptend_q0002_51', 'ptend_q0002_52', 'ptend_q0002_53', 'ptend_q0002_54', 'ptend_q0002_55', 'ptend_q0002_56', 'ptend_q0002_57', 'ptend_q0002_58', 'ptend_q0002_59',],\n    'ptend_q0003': ['ptend_q0003_0', 'ptend_q0003_1', 'ptend_q0003_2', 'ptend_q0003_3', 'ptend_q0003_4', 'ptend_q0003_5', 'ptend_q0003_6', 'ptend_q0003_7', 'ptend_q0003_8', 'ptend_q0003_9', 'ptend_q0003_10', 'ptend_q0003_11', 'ptend_q0003_12', 'ptend_q0003_13', 'ptend_q0003_14', 'ptend_q0003_15', 'ptend_q0003_16', 'ptend_q0003_17', 'ptend_q0003_18', 'ptend_q0003_19', 'ptend_q0003_20', 'ptend_q0003_21', 'ptend_q0003_22', 'ptend_q0003_23', 'ptend_q0003_24', 'ptend_q0003_25', 'ptend_q0003_26', 'ptend_q0003_27', 'ptend_q0003_28', 'ptend_q0003_29', 'ptend_q0003_30', 'ptend_q0003_31', 'ptend_q0003_32', 'ptend_q0003_33', 'ptend_q0003_34', 'ptend_q0003_35', 'ptend_q0003_36', 'ptend_q0003_37', 'ptend_q0003_38', 'ptend_q0003_39', 'ptend_q0003_40', 'ptend_q0003_41', 'ptend_q0003_42', 'ptend_q0003_43', 'ptend_q0003_44', 'ptend_q0003_45', 'ptend_q0003_46', 'ptend_q0003_47', 'ptend_q0003_48', 'ptend_q0003_49', 'ptend_q0003_50', 'ptend_q0003_51', 'ptend_q0003_52', 'ptend_q0003_53', 'ptend_q0003_54', 'ptend_q0003_55', 'ptend_q0003_56', 'ptend_q0003_57', 'ptend_q0003_58', 'ptend_q0003_59',],\n    'cam_out': ['cam_out_NETSW', 'cam_out_FLWDS', 'cam_out_PRECSC', 'cam_out_PRECC', 'cam_out_SOLS', 'cam_out_SOLL', 'cam_out_SOLSD', 'cam_out_SOLLD'],\n}\nprint('target_groups:', target_groups.keys())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport glob\nimport matplotlib.pyplot as plt\nimport seaborn as sns","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# read parquet","metadata":{}},{"cell_type":"code","source":"df = pd.concat([\n    pd.read_parquet(\n        fn,\n    ) for fn in sorted(glob.glob(f'/kaggle/working/leap-atmospheric-physics-ai-climsim/train_*.parquet'))\n], axis=0).reset_index().sort_index()\ndf['id'] = df['sample_id'].str.extract('(\\d+)').astype(int)\ndf.drop('sample_id', axis=1, inplace=True)\ndf","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# plot first 1000 rows","metadata":{}},{"cell_type":"code","source":"n_samples = 10000\nn_samples_reduced = int(n_samples / 10)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# feature_groups","metadata":{}},{"cell_type":"code","source":"for group, feats in feature_groups.items():\n    fig, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=(24, 6))\n    \n    sns.heatmap(df[feats].head(n_samples).T, cmap='viridis', ax=ax1)\n    ax1.set_xlabel('Dataset Index')\n    ax1.set_ylabel('Features')\n    ax1.set_title(f'{group} - Heatmap (index < {n_samples})')\n    ax1.set_xticks(range(0, n_samples, int(n_samples/10)))\n    ax1.set_xticklabels(range(0, n_samples, int(n_samples/10)))\n    \n    sns.heatmap(df[feats].head(n_samples_reduced).T, cmap='viridis', ax=ax2)\n    ax2.set_xlabel('Dataset Index')\n    ax2.set_ylabel('Features')\n    ax2.set_title(f'{group} - Heatmap (index < {n_samples_reduced})')\n    ax2.set_xticks(range(0, n_samples_reduced, int(n_samples_reduced/10)))\n    ax2.set_xticklabels(range(0, n_samples_reduced, int(n_samples_reduced/10)))\n    \n    sns.boxplot(data=df[feats].head(n_samples), ax=ax3)\n    ax3.set_xticklabels(ax3.get_xticklabels(), rotation=90)\n    ax3.set_title(f'{group} - Box Plot')\n    ax3.set_xlabel('Features')\n    ax3.set_ylabel('Value')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport seaborn as sns\n\ncorr = {}\n\nfor group, feats in feature_groups.items():\n    data_sample = df[feats].head(n_samples)\n    data_sample_reduced = df[feats].head(n_samples_reduced)\n    \n    num_features = len(feats)\n    if group in ['cam', 'pbuf']:\n        # 'cam' または 'pbuf' グループの各特徴量ごとにプロットを作成\n        for idx, feat in enumerate(feats):\n            fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6))\n            ax1.plot(data_sample.index, data_sample[feat], label=f'{feat} Full')\n            ax1.set_title(f'Time Series Analysis for {feat} (Index < {n_samples})')\n            ax1.set_xlabel('Time Index')\n            ax1.set_ylabel('Value')\n            ax1.legend()\n            ax1.grid(True)\n\n            ax2.plot(data_sample_reduced.index, data_sample_reduced[feat], label=f'{feat} Reduced')\n            ax2.set_title(f'Time Series Analysis for {feat} (Index < {n_samples_reduced})')\n            ax2.set_xlabel('Time Index')\n            ax2.set_ylabel('Value')\n            ax2.legend()\n            ax2.grid(True)\n\n            plt.tight_layout()\n            plt.show()\n            \n            corr[feat] = data_sample[feat]\n        \n    else:\n        fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6))\n        ax1.plot(data_sample.mean(axis=1), label='Mean Full', color='blue')\n        ax1.plot(data_sample.min(axis=1), label='Min Full', color='green')\n        ax1.plot(data_sample.max(axis=1), label='Max Full', color='red')\n        ax1.set_title(f'{group} Analysis (Index < {n_samples})')\n        ax1.set_xlabel('Time Index')\n        ax1.set_ylabel('Value')\n        ax1.legend()\n        ax1.grid(True)\n\n        ax2.plot(data_sample_reduced.mean(axis=1), label='Mean Reduced', color='blue')\n        ax2.plot(data_sample_reduced.min(axis=1), label='Min Reduced', color='green')\n        ax2.plot(data_sample_reduced.max(axis=1), label='Max Reduced', color='red')\n        ax2.set_title(f'{group} Analysis (Index < {n_samples_reduced})')\n        ax2.set_xlabel('Time Index')\n        ax2.set_ylabel('Value')\n        ax2.legend()\n        ax2.grid(True)\n\n        plt.tight_layout()\n        plt.show()\n\n        corr[group+'_mean'] = data_sample.mean(axis=1)\n        corr[group+'_min'] = data_sample.min(axis=1)\n        corr[group+'_max'] = data_sample.max(axis=1)\n\ncorr_features = pd.DataFrame.from_dict(corr)\n# 相関係数の計算\ncorr_matrix = corr_features.corr()\n\n# 相関行列から上三角部分を抽出（重複を避けるため）\nmask = np.triu(np.ones_like(corr_matrix, dtype=bool))\n\n# 相関行列のヒートマップ\nplt.figure(figsize=(16, 16))\nsns.heatmap(corr_matrix, mask=mask, cmap='coolwarm', cbar=True, linewidths=.5, square=True)\nplt.title('Correlation Matrix')\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# target_groups","metadata":{}},{"cell_type":"code","source":"for group, feats in target_groups.items():\n    fig, (ax1, ax2, ax3) = plt.subplots(1, 3, figsize=(24, 6))\n    \n    sns.heatmap(df[feats].head(n_samples).T, cmap='viridis', ax=ax1)\n    ax1.set_xlabel('Dataset Index')\n    ax1.set_ylabel('Features')\n    ax1.set_title(f'{group} - Heatmap (index < {n_samples})')\n    ax1.set_xticks(range(0, n_samples, int(n_samples/10)))\n    ax1.set_xticklabels(range(0, n_samples, int(n_samples/10)))\n    \n    sns.heatmap(df[feats].head(n_samples_reduced).T, cmap='viridis', ax=ax2)\n    ax2.set_xlabel('Dataset Index')\n    ax2.set_ylabel('Features')\n    ax2.set_title(f'{group} - Heatmap (index < {n_samples_reduced})')\n    ax2.set_xticks(range(0, n_samples_reduced, int(n_samples_reduced/10)))\n    ax2.set_xticklabels(range(0, n_samples_reduced, int(n_samples_reduced/10)))\n    \n    sns.boxplot(data=df[feats].head(n_samples), ax=ax3)\n    ax3.set_xticklabels(ax3.get_xticklabels(), rotation=90)\n    ax3.set_title(f'{group} - Box Plot')\n    ax3.set_xlabel('Features')\n    ax3.set_ylabel('Value')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport seaborn as sns\n\ncorr = {}\n\nfor group, feats in target_groups.items():\n    data_sample = df[feats].head(n_samples)\n    data_sample_reduced = df[feats].head(n_samples_reduced)\n    \n    num_features = len(feats)\n    if group in ['cam_out']:\n        for idx, feat in enumerate(feats):\n            fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6))\n            ax1.plot(data_sample.index, data_sample[feat], label=f'{feat} Full')\n            ax1.set_title(f'Time Series Analysis for {feat} (Index < {n_samples})')\n            ax1.set_xlabel('Time Index')\n            ax1.set_ylabel('Value')\n            ax1.legend()\n            ax1.grid(True)\n\n            ax2.plot(data_sample_reduced.index, data_sample_reduced[feat], label=f'{feat} Reduced')\n            ax2.set_title(f'Time Series Analysis for {feat} (Index < {n_samples_reduced})')\n            ax2.set_xlabel('Time Index')\n            ax2.set_ylabel('Value')\n            ax2.legend()\n            ax2.grid(True)\n\n            plt.tight_layout()\n            plt.show()\n            \n            corr[feat] = data_sample[feat]\n        \n    else:\n        fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6))\n        ax1.plot(data_sample.mean(axis=1), label='Mean Full', color='blue')\n        ax1.plot(data_sample.min(axis=1), label='Min Full', color='green')\n        ax1.plot(data_sample.max(axis=1), label='Max Full', color='red')\n        ax1.set_title(f'{group} Analysis (Index < {n_samples})')\n        ax1.set_xlabel('Time Index')\n        ax1.set_ylabel('Value')\n        ax1.legend()\n        ax1.grid(True)\n\n        ax2.plot(data_sample_reduced.mean(axis=1), label='Mean Reduced', color='blue')\n        ax2.plot(data_sample_reduced.min(axis=1), label='Min Reduced', color='green')\n        ax2.plot(data_sample_reduced.max(axis=1), label='Max Reduced', color='red')\n        ax2.set_title(f'{group} Analysis (Index < {n_samples_reduced})')\n        ax2.set_xlabel('Time Index')\n        ax2.set_ylabel('Value')\n        ax2.legend()\n        ax2.grid(True)\n\n        plt.tight_layout()\n        plt.show()\n\n        corr[group+'_mean'] = data_sample.mean(axis=1)\n        corr[group+'_min'] = data_sample.min(axis=1)\n        corr[group+'_max'] = data_sample.max(axis=1)\n\ncorr_targets = pd.DataFrame.from_dict(corr)\n# 相関係数の計算\ncorr_matrix = corr_targets.corr()\n\n# 相関行列から上三角部分を抽出（重複を避けるため）\nmask = np.triu(np.ones_like(corr_matrix, dtype=bool))\n\n# 相関行列のヒートマップ\nplt.figure(figsize=(16, 16))\nsns.heatmap(corr_matrix, mask=mask, cmap='coolwarm', cbar=True, linewidths=.5, square=True)\nplt.title('Correlation Matrix')\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Correlation between Features and Targets","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport numpy as np\n\n# 相関係数を計算し、結果を格納するためのデータフレームを準備\ncorr_values = []\n\n# corr_featuresとcorr_targetsの全ての組み合わせで相関係数を計算\nfor feature_col in corr_features.columns:\n    for target_col in corr_targets.columns:\n        if np.all(np.isfinite(corr_features[feature_col])) and np.all(np.isfinite(corr_targets[target_col])):\n            corr_coef = np.corrcoef(corr_features[feature_col], corr_targets[target_col])[0, 1]\n            corr_values.append((feature_col, target_col, corr_coef))\n        else:\n            corr_values.append((feature_col, target_col, np.nan))  # Handle non-finite values\n\n# 相関データをデータフレームに変換\ncorr_df = pd.DataFrame(corr_values, columns=['Feature', 'Target', 'Correlation'])\n\n# 散布図のプロット\nplt.figure(figsize=(20, 12))\nax = plt.gca()  # 現在の軸を取得\nscatter = sns.scatterplot(\n    data=corr_df,\n    x='Feature',\n    y='Target',\n    size=np.abs(corr_df['Correlation']),  # Use absolute value for sizes\n    hue='Correlation',\n    sizes=(100, 1000),  # Adjust the range of sizes based on absolute correlation\n    palette='coolwarm',  # Red for positive, blue for negative\n    hue_norm=(-1, 1),  # Normalize hue to the correlation range\n    legend=False,  # Disable the default legend\n    alpha=0.6,  # Adjust transparency\n    ax=ax  # Specify the axis object explicitly\n)\n\n# カスタムカラーバーを追加\nnorm = plt.Normalize(-1, 1)\nsm = plt.cm.ScalarMappable(cmap=\"coolwarm\", norm=norm)\nsm.set_array([])\ncbar = plt.colorbar(sm, ax=ax, ticks=np.linspace(-1, 1, 5))  # Specify the axis for the colorbar\ncbar.set_label('Correlation Coefficient')\n\nplt.title('Correlation between Features and Targets')\nplt.xticks(rotation=90)\nplt.grid(True)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Clustering","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.preprocessing import MinMaxScaler, StandardScaler\nfrom sklearn.decomposition import PCA\nfrom sklearn.cluster import KMeans, DBSCAN, AgglomerativeClustering\nfrom sklearn.metrics import silhouette_score\nfrom sklearn.manifold import TSNE\n\n# データのスケーリング\nscaler = MinMaxScaler()\nscaled_features = scaler.fit_transform(corr_features)\n\n# PCAによる次元削減\npca = PCA(n_components=2)\nreduced_features = pca.fit_transform(scaled_features)\n\n# シルエットスコアを用いて最適なクラスタ数を決定\nrange_n_clusters = list(range(2, 10))\nsilhouette_avg_scores = []\n\nfor n_clusters in range_n_clusters:\n    kmeans = KMeans(n_clusters=n_clusters, random_state=42, n_init=10)\n    cluster_labels = kmeans.fit_predict(reduced_features)\n    silhouette_avg = silhouette_score(reduced_features, cluster_labels)\n    silhouette_avg_scores.append(silhouette_avg)\n\n# 最適なクラスタ数を選択\noptimal_clusters = range_n_clusters[np.argmax(silhouette_avg_scores)]\n\n# K-means クラスタリング\nkmeans = KMeans(n_clusters=optimal_clusters, random_state=42, n_init=10)\nkmeans_labels = kmeans.fit_predict(reduced_features)\n\n# DBSCAN クラスタリング\ndbscan = DBSCAN(eps=0.1, min_samples=5)\ndbscan_labels = dbscan.fit_predict(reduced_features)\n\n# Agglomerative Clustering\nagglo = AgglomerativeClustering(n_clusters=optimal_clusters, linkage='ward')\nagglo_labels = agglo.fit_predict(reduced_features)\n\n# 可視化\nfig, axs = plt.subplots(1, 3, figsize=(18, 6))\n\n# K-means の結果\naxs[0].scatter(reduced_features[:, 0], reduced_features[:, 1], c=kmeans_labels, cmap='viridis', edgecolor='k', alpha=0.6)\naxs[0].set_title('K-means Clustering (n_clusters=' + str(optimal_clusters) + ')')\naxs[0].set_xlabel('PCA Dimension 1')\naxs[0].set_ylabel('PCA Dimension 2')\n\n# DBSCAN の結果\naxs[1].scatter(reduced_features[:, 0], reduced_features[:, 1], c=dbscan_labels, cmap='plasma', edgecolor='k', alpha=0.6)\naxs[1].set_title('DBSCAN Clustering')\naxs[1].set_xlabel('PCA Dimension 1')\naxs[1].set_ylabel('PCA Dimension 2')\n\n# Agglomerative Clustering の結果\naxs[2].scatter(reduced_features[:, 0], reduced_features[:, 1], c=agglo_labels, cmap='inferno', edgecolor='k', alpha=0.6)\naxs[2].set_title('Agglomerative Clustering')\naxs[2].set_xlabel('PCA Dimension 1')\naxs[2].set_ylabel('PCA Dimension 2')\n\nplt.tight_layout()\nplt.show()\n\n# シルエットスコアのプロット\nplt.figure(figsize=(10, 6))\nplt.plot(range_n_clusters, silhouette_avg_scores, marker='o')\nplt.title('Silhouette Score to Determine Optimal Cluster Number')\nplt.xlabel('Number of Clusters')\nplt.ylabel('Silhouette Score')\nplt.grid(True)\nplt.show()\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import tensorflow as tf\nfrom tensorflow.keras.models import Model\nfrom tensorflow.keras.layers import Input, Dense\nfrom tensorflow.keras.optimizers import Adam\nfrom sklearn.metrics import silhouette_score\nfrom sklearn.cluster import KMeans\nimport matplotlib.pyplot as plt\nimport numpy as np\n\ntf.keras.utils.set_random_seed(42)\n\n# モデル定義と訓練\ninput_layer = Input(shape=(scaled_features.shape[1],))\nencoded = Dense(64, activation='relu')(input_layer)\nencoded = Dense(32, activation='relu')(encoded)\ndecoded = Dense(64, activation='relu')(encoded)\ndecoded = Dense(scaled_features.shape[1], activation='sigmoid')(decoded)\n\nautoencoder = Model(input_layer, decoded)\nautoencoder.compile(optimizer='adam', loss='mse')\nencoder = Model(input_layer, encoded)\nautoencoder.fit(scaled_features, scaled_features, epochs=50, batch_size=256, shuffle=True, verbose=0)\n\n# エンコードされた特徴を抽出\nencoded_features = encoder.predict(scaled_features)\n\n# シルエットスコアによる最適なクラスタ数の決定\nrange_n_clusters = list(range(2, 10))\nsilhouette_avg_scores = []\n\nfor n_clusters in range_n_clusters:\n    kmeans = KMeans(n_clusters=n_clusters, random_state=42, n_init=10)\n    cluster_labels = kmeans.fit_predict(encoded_features)\n    silhouette_avg = silhouette_score(encoded_features, cluster_labels)\n    silhouette_avg_scores.append(silhouette_avg)\n\noptimal_clusters = range_n_clusters[np.argmax(silhouette_avg_scores)]\nprint(\"Optimal number of clusters: \", optimal_clusters)\n\n# クラスタリングと可視化\nkmeans = KMeans(n_clusters=optimal_clusters, random_state=42, n_init=10)\nae_labels = kmeans.fit_predict(encoded_features)\n\nplt.figure(figsize=(10, 6))\nplt.scatter(encoded_features[:, 0], encoded_features[:, 1], c=ae_labels, cmap='viridis', alpha=0.6, edgecolors='k')\nplt.title('Clustering with Autoencoder Features')\nplt.xlabel('Encoded Dimension 1')\nplt.ylabel('Encoded Dimension 2')\nplt.colorbar(label='Cluster Label')\nplt.show()\n\n# シルエットスコアのプロット\nplt.figure(figsize=(10, 6))\nplt.plot(range_n_clusters, silhouette_avg_scores, marker='o')\nplt.title('Silhouette Scores for Optimal Cluster Number')\nplt.xlabel('Number of Clusters')\nplt.ylabel('Silhouette Score')\nplt.grid(True)\nplt.show()\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}