{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":2},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython2","version":"2.7.6"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":59094,"databundleVersionId":7010844,"sourceType":"competition"},{"sourceId":6947995,"sourceType":"datasetVersion","datasetId":3923807}],"dockerImageVersionId":30558,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"Topological Data Analysis (TDA) using persistent homology features generated from 3-D molecular shapes as described in \"Shape is (almost) all!: Persistent homology features (PHFs) are an information rich input for efficient molecular machine learning\" [arXiv:2304.07554 \\[cs.LG\\]](https://arxiv.org/abs/2304.07554)\n- RDKit to convert the SMILES data to 3-D point clouds\n- PyWGCNA to group genes into modules\n- Giotto-tda to generate the persistent homology features\n","metadata":{}},{"cell_type":"code","source":"import os\nimport joblib\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as mpatches\nfrom mpl_toolkits.mplot3d import Axes3D, proj3d\nimport random\nimport seaborn as sns\nimport sys\nfrom typing import List\n\nfrom sklearn.preprocessing import OneHotEncoder, OrdinalEncoder\nfrom sklearn.decomposition import TruncatedSVD\nfrom sklearn.decomposition import PCA\nfrom sklearn.compose import ColumnTransformer\nfrom sklearn.pipeline import make_pipeline\nfrom sklearn.linear_model import Ridge\nfrom sklearn.ensemble import RandomForestRegressor\nfrom sklearn.svm import LinearSVR\nfrom sklearn.metrics import mean_squared_error, r2_score\nfrom sklearn.multioutput import MultiOutputRegressor\nfrom sklearn.model_selection import KFold, train_test_split\nfrom sklearn.neural_network import MLPRegressor\nfrom sklearn.preprocessing import StandardScaler\nfrom sklearn.pipeline import make_pipeline, make_union\nfrom sklearn import linear_model\nfrom sklearn.decomposition import TruncatedSVD\nfrom sklearn.compose import ColumnTransformer\nfrom sklearn.linear_model import Ridge\n\nfrom IPython.display import Javascript, SVG  # Restrict height of output cell.\n\nimport plotly\nimport matplotlib.image as mpimg\nimport io\n\nimport warnings\nimport pickle\nfrom xgboost import XGBRegressor\n","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2023-11-18T00:21:53.895634375Z","start_time":"2023-11-18T00:21:53.489225939Z"},"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if os.environ.get('KAGGLE_KERNEL_RUN_TYPE', 'Localhost') != 'Localhost':\n    !pip install rdkit\n    !pip install giotto-tda\n    !pip install kaleido\n","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2023-11-18T00:21:53.897932465Z","start_time":"2023-11-18T00:21:53.888276461Z"},"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# RDkit\nfrom rdkit import rdBase, Chem\nfrom rdkit.Chem import AllChem, AddHs, Descriptors, Draw\nfrom rdkit.Chem.rdmolops import GetAdjacencyMatrix, GetDistanceMatrix, Get3DDistanceMatrix\nfrom rdkit.Chem.Draw import IPythonConsole, rdMolDraw2D\nfrom rdkit import RDLogger\n\nfrom gtda.homology import VietorisRipsPersistence\nfrom gtda.plotting import plot_diagram, plot_point_cloud\nfrom gtda.diagrams import Amplitude, NumberOfPoints, PersistenceEntropy\n\n","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2023-11-18T00:21:54.008596486Z","start_time":"2023-11-18T00:21:53.888988812Z"},"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"file_path = '../data/'\nif os.environ.get('KAGGLE_KERNEL_RUN_TYPE', 'Localhost') != 'Localhost':\n    file_path = '/kaggle/input/open-problems-single-cell-perturbations/'\n\nid_map = pd.read_csv(file_path+'id_map.csv', index_col='id')\nde_train = pd.read_parquet(file_path+'de_train.parquet')\ndisplay(de_train)\n","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2023-11-18T00:22:01.505575240Z","start_time":"2023-11-18T00:21:54.014076733Z"},"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sample_submission = pd.read_csv(file_path+'sample_submission.csv')\ngenes = de_train.columns[5:] # 18211 genes\n\n# creat an sm_name to SMILES dictionary for the prediction phase\nsm_name_to_SMILES_dict = pd.Series(de_train.SMILES.values, index=de_train.sm_name).to_dict()\ndisplay(sm_name_to_SMILES_dict)\n","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2023-11-18T00:22:04.858061763Z","start_time":"2023-11-18T00:22:01.514012168Z"},"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if os.environ.get('KAGGLE_KERNEL_RUN_TYPE', 'Localhost') != 'Localhost':\n    black_lst = pickle.load(open('/kaggle/input/opscp001/black_lst.pkl', 'rb'))\n    darkgrey_lst = pickle.load(open('/kaggle/input/opscp001/darkgrey_lst.pkl', 'rb'))\n    silver_lst = pickle.load(open('/kaggle/input/opscp001/silver_lst.pkl', 'rb'))\nelse:\n    black_lst = pickle.load(open('black_lst.pkl', 'rb'))\n    darkgrey_lst = pickle.load(open('darkgrey_lst.pkl', 'rb'))\n    silver_lst = pickle.load(open('silver_lst.pkl', 'rb'))\n    # import PyWGCNA\n    #\n    # geneExp = de_train.iloc[:, 5:]\n    # pyWGCNA_CC = PyWGCNA.WGCNA(name=' Open Problems – Single-Cell Perturbations',\n    #                            geneExp=geneExp,\n    #                            outputPath='',\n    #                            save=False)\n    # pyWGCNA_CC.geneExpr.to_df().head(5)\n    # pyWGCNA_CC.preprocess()\n    # pyWGCNA_CC.findModules()\n    #\n    # modules = pyWGCNA_CC.getModuleName()\n    # display(modules)\n    #\n    # black_lst = list(pyWGCNA_CC.getGeneModule('black')['black'].index)\n    # darkgrey_lst = list(pyWGCNA_CC.getGeneModule('darkgrey')['darkgrey'].index)\n    # silver_lst = list(pyWGCNA_CC.getGeneModule('silver')['silver'].index)\n    # assert(set(black_lst).isdisjoint(darkgrey_lst))\n    # assert(set(darkgrey_lst).isdisjoint(silver_lst))\n    # assert(set(silver_lst).isdisjoint(black_lst))\n    # assert(not set(black_lst).isdisjoint(black_lst))\n    # pickle.dump(black_lst,open('black_lst.pkl','wb'))\n    # pickle.dump(darkgrey_lst,open('darkgrey_lst.pkl','wb'))\n    # pickle.dump(silver_lst,open('silver_lst.pkl','wb'))\n\nprint(f\"len(black_lst): {len(black_lst)}, len(darkgrey_lst): {len(darkgrey_lst)}, len(silver_lst): {len(silver_lst)}\")\n","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2023-11-18T00:22:04.937979641Z","start_time":"2023-11-18T00:22:04.867631632Z"},"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def mol_from_smiles(smiles):\n    with warnings.catch_warnings():\n        warnings.simplefilter(\"ignore\")\n        RDLogger.DisableLog('rdApp.*')\n        mol = Chem.MolFromSmiles(smiles)\n        mol = Chem.AddHs(mol)\n        return mol\n\ndef point_cloud_from_smiles(smiles):\n    with warnings.catch_warnings():\n        warnings.simplefilter(\"ignore\")\n        RDLogger.DisableLog('rdApp.*')\n        # Generate a 3-D structure from smiles\n        mol = Chem.MolFromSmiles(smiles)\n        mol = Chem.AddHs(mol)\n        p = AllChem.ETKDG()\n        param = AllChem.ETKDGv3()\n        param.randomSeed = 42 # Distance geometry で立体配座を生成\n        status = AllChem.EmbedMolecule(mol, param)\n        status = AllChem.UFFOptimizeMolecule(mol)\n        conformer = mol.GetConformer()\n        coordinates = conformer.GetPositions()\n        RDLogger.EnableLog('rdApp.*')\n        coordinates = np.array(coordinates)\n        return coordinates\n\n\ndef generate_tda_features_v3(df, metrics):\n    if metrics == 'metric_lst_1':\n        metric_lst = [\"bottleneck\", \"wasserstein\", \"persistence_image\"]\n    elif metrics == 'metric_lst_2':\n        metric_lst = ['bottleneck', 'wasserstein', 'betti', 'landscape', 'silhouette', 'heat', 'persistence_image']\n    else:\n        metric_lst = ['landscape']\n\n    homology_dimensions = [0, 1, 2]\n    metrics = [ {\"metric\": metric} for metric in metric_lst ]\n    feature_union = make_union(\n        PersistenceEntropy(),\n        NumberOfPoints(n_jobs=-1),\n        *[Amplitude(**metric, n_jobs=-1) for metric in metrics]\n    )\n\n    steps = [VietorisRipsPersistence(metric=\"euclidean\",\n                                     homology_dimensions=homology_dimensions,\n                                     n_jobs=-1,\n                                     collapse_edges=True, ),\n             feature_union,\n             ]\n    pipeline = make_pipeline(*steps)\n\n    if 'SMILES' not in df.columns:\n        df['SMILES'] = df['sm_name'].map(sm_name_to_SMILES_dict)\n    smiles_lst = df['SMILES'].values.tolist()\n    coords_lst = joblib.Parallel(n_jobs=-1)(joblib.delayed(point_cloud_from_smiles)(x) for x in smiles_lst)\n    tda_features = pipeline.fit_transform(coords_lst)\n    return tda_features\n\n","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2023-11-18T00:29:36.246306241Z","start_time":"2023-11-18T00:29:36.205892957Z"},"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"NUM_MOLS = 5\nmol_name_lst = []\nmol_lst = []\nmol_img_lst = []\nmol_pt_cld_lst = []\nmols = list(sm_name_to_SMILES_dict.keys())\nfor i in range(NUM_MOLS):\n    mol_name_lst.append(mols[i])\n    smiles = sm_name_to_SMILES_dict[mols[i]]\n    m = mol_from_smiles(smiles)\n    mol_lst.append(m)\n    img = Draw.MolToImage(m)\n    mol_img_lst.append(img)\n    pt_cld = point_cloud_from_smiles(smiles)\n    mol_pt_cld_lst.append(pt_cld)\nhomology_dimensions = [0, 1, 2]\nVR = VietorisRipsPersistence(metric=\"euclidean\",\n                             homology_dimensions=homology_dimensions,\n                             n_jobs=-1,\n                             collapse_edges=True, )\nmol_diagrams_lst = VR.fit_transform(mol_pt_cld_lst)\n# metric_lst = ['bottleneck', 'wasserstein', 'betti', 'landscape', 'silhouette', 'heat', 'persistence_image']\nmetric_lst = [\"bottleneck\", \"wasserstein\", \"persistence_image\"]\nmetrics = [ {\"metric\": metric} for metric in metric_lst ]\nfeature_union = make_union(\n    PersistenceEntropy(),\n    NumberOfPoints(n_jobs=-1),\n    *[Amplitude(**metric, n_jobs=-1) for metric in metrics]\n)\nmol_pe_feature_lst = feature_union.fit_transform(mol_diagrams_lst)\n\nf, axarr = plt.subplots(NUM_MOLS, 3, figsize=(20, 25))\nfor idx_1 in range(NUM_MOLS):\n    axarr[idx_1, 0].imshow(mol_img_lst[idx_1])\n    axarr[idx_1, 0].set_title(mol_name_lst[idx_1])\n    axarr[idx_1, 0].axis('off')\n\n    bytes_data = plotly.io.to_image(plot_point_cloud(mol_diagrams_lst[idx_1]), 'png')\n    fp = io.BytesIO(bytes_data)\n    with fp:\n        img = mpimg.imread(fp, format='png')\n    axarr[idx_1, 1].imshow(img)\n    axarr[idx_1, 1].set_title(mol_name_lst[idx_1])\n    axarr[idx_1, 1].axis('off')\n\n    bytes_data = plotly.io.to_image(plot_diagram(mol_diagrams_lst[idx_1]), 'png')\n    fp = io.BytesIO(bytes_data)\n    with fp:\n        img = mpimg.imread(fp, format='png')\n    axarr[idx_1, 2].imshow(img)\n    axarr[idx_1, 2].set_title('Persistence Diagram')\n    axarr[idx_1, 2].axis('off')\n","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2023-11-18T00:29:52.131430459Z","start_time":"2023-11-18T00:29:36.822134134Z"},"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_predict_tda_v3(de_all_df, df_test, model, n_components, features, metrics, random_seed=42):\n    np.random.seed(random_seed)\n    oh = OneHotEncoder(sparse_output=False)\n    oh.fit(de_all_df[features])\n\n    X_train_all_oh = oh.transform(de_all_df[features])\n    tda_features = generate_tda_features_v3(de_all_df, metrics=metrics)\n    X_train_all_oh = np.concatenate((X_train_all_oh, tda_features), axis=1)\n\n    svd_black = TruncatedSVD(n_components=n_components)\n    # svd_darkgrey = TruncatedSVD(n_components=n_components)\n    svd_silver = TruncatedSVD(n_components=n_components)\n    y_train_all = de_all_df.iloc[:, 5:]\n    y_train_all_black = svd_black.fit_transform(y_train_all[black_lst])\n    y_train_all_silver = svd_silver.fit_transform(y_train_all[silver_lst])\n    y_train_all = np.concatenate([y_train_all_black, y_train_all[darkgrey_lst], y_train_all_silver], 1)\n\n    model.fit(X_train_all_oh, y_train_all)\n\n    # Predict\n    X_test_all_oh = oh.transform(df_test[features])\n    tda_features = generate_tda_features_v3(df_test, metrics=metrics)\n    X_test_all_oh = np.concatenate((X_test_all_oh, tda_features), axis=1)\n\n    y_preds = pd.DataFrame(columns=genes,\n                           index=df_test.index)\n    y_module_preds = model.predict(X_test_all_oh)\n    y_preds.loc[:, black_lst] = svd_black.inverse_transform(y_module_preds[:, :n_components])\n    y_preds.loc[:, darkgrey_lst] = y_module_preds[:, n_components:-n_components]\n    y_preds.loc[:, silver_lst] = svd_silver.inverse_transform(y_module_preds[:, -n_components:])\n    return y_preds\n\n","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2023-11-18T00:33:33.504053245Z","start_time":"2023-11-18T00:33:33.426787247Z"},"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = Ridge(alpha=1.5, fit_intercept=True)\nfeatures = [\"sm_name\"]\nmetrics = 'metric_lst_2'\nn_components=40\ny_preds = train_predict_tda_v3(de_train, id_map, model, n_components, features=features, metrics=metrics, random_seed=42)\nsubmission = pd.DataFrame(y_preds, columns=genes, index=id_map.index)\n\ndisplay(submission)\nsubmission.to_csv('submission.csv')\n","metadata":{"collapsed":false,"ExecuteTime":{"end_time":"2023-11-18T00:34:15.248697281Z","start_time":"2023-11-18T00:33:34.393506923Z"},"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"collapsed":false,"is_executing":true,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"collapsed":false,"jupyter":{"outputs_hidden":false}},"execution_count":null,"outputs":[]}]}