{"cells":[{"cell_type":"code","execution_count":null,"id":"af118806","metadata":{"language":"python"},"outputs":[],"source":"from pathlib import Path\n\nfrom IPython.display import Image, Markdown, display\n\n\n\nSTATIC_IMAGE_DIR = Path(\"/kaggle/input/notebook-static-images-public\")\n\ndisplay(Markdown(\"## Saved analysis figures\"))\n\nfor image_path in sorted(STATIC_IMAGE_DIR.glob(\"integration-*.png\")):\n\n    display(Markdown(f\"**{image_path.stem}**\"))\n\n    display(Image(filename=str(image_path)))"},{"cell_type":"markdown","id":"13d0758d-7b69-4f84-bf66-1bd6c8d18c56","metadata":{"jp-MarkdownHeadingCollapsed":true},"source":"# Data merge : Multiome(scRNA + scATAC)"},{"cell_type":"code","execution_count":null,"id":"4fea546a-0419-4610-8c9f-181ac30d8e47","metadata":{},"outputs":[],"source":"import pandas as pd\n\ninput_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/train/train_multi_inputs.h5\"\ntarget_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/train/train_multi_targets.h5\"\n\ntrain_multi_inputs = pd.read_hdf(input_path, key=\"train_multi_inputs\")\ntrain_multi_targets = pd.read_hdf(target_path, key=\"train_multi_targets\")\n\nprint(train_multi_inputs.shape)\nprint(train_multi_targets.shape)"},{"cell_type":"code","execution_count":null,"id":"469dffd0-6291-4230-a662-6a7908bc22df","metadata":{},"outputs":[],"source":"import pandas as pd\nimport numpy as np\nimport scipy.sparse as sp\nimport anndata as ad\nimport pyranges as pr\n\nbase = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/baseline\"\n\npromoter_path = f\"{base}/promoter.bed\"\nenhancer_path = f\"{base}/enhancer.bed\"\ngtf_path = f\"{base}/gencode.v44.annotation.gtf.gz\"\n\ngtf = pd.read_csv(\n    gtf_path,\n    sep=\"\\t\",\n    comment=\"#\",\n    header=None,\n    names=[\n        \"Chromosome\",\n        \"Source\",\n        \"Feature\",\n        \"Start\",\n        \"End\",\n        \"Score\",\n        \"Strand\",\n        \"Frame\",\n        \"Attribute\"\n    ]\n)\n\ngtf_gene = gtf[gtf[\"Feature\"] == \"gene\"].copy()\n\ngtf_gene[\"gene_id\"] = gtf_gene[\"Attribute\"].str.extract(\n    r'gene_id \"([^\"]+)\"'\n)[0]\n\ngtf_gene[\"gene_name\"] = gtf_gene[\"Attribute\"].str.extract(\n    r'gene_name \"([^\"]+)\"'\n)[0]\n\ngtf_gene[\"gene_id\"] = gtf_gene[\"gene_id\"].str.replace(\n    r\"\\.\\d+$\", \"\", regex=True\n)\n\ntarget = train_multi_targets.copy()\ntarget.index = target.index.astype(str)\n\ntarget_ensg = pd.Index(\n    target.columns.astype(str)\n).str.replace(\n    r\"\\.\\d+$\", \"\", regex=True\n)\n\nensg_to_symbol = (\n    gtf_gene[\n        [\"gene_id\", \"gene_name\"]\n    ]\n    .dropna()\n    .drop_duplicates(\"gene_id\")\n    .set_index(\"gene_id\")[\"gene_name\"]\n    .to_dict()\n)\n\ntarget_symbols = target_ensg.map(ensg_to_symbol)\n\nvalid_target = target_symbols.notna()\n\ntarget = target.loc[:, valid_target]\ntarget.columns = target_symbols[valid_target]\n\ntarget = target.loc[:, ~target.columns.duplicated()]\n\ngene_order = pd.Index(target.columns.astype(str))\n\ngtf_gene = gtf_gene[\n    gtf_gene[\"gene_id\"].isin(\n        target_ensg[valid_target]\n    )\n].copy()\n\ngene_tss = gtf_gene[\n    [\n        \"Chromosome\",\n        \"Start\",\n        \"End\",\n        \"Strand\",\n        \"gene_id\",\n        \"gene_name\"\n    ]\n].copy()\n\ngene_tss[\"Start\"] = np.where(\n    gene_tss[\"Strand\"] == \"+\",\n    gene_tss[\"Start\"] - 1,\n    gene_tss[\"End\"] - 1\n)\n\ngene_tss[\"End\"] = gene_tss[\"Start\"] + 1\n\ngene_tss_pr = pr.PyRanges(gene_tss)\n\npromoter = pd.read_csv(\n    promoter_path,\n    sep=\"\\t\",\n    header=None,\n    names=[\n        \"Chromosome\",\n        \"Start\",\n        \"End\",\n        \"ID1\",\n        \"ID2\",\n        \"Annotation\"\n    ]\n)\n\nenhancer = pd.read_csv(\n    enhancer_path,\n    sep=\"\\t\",\n    header=None,\n    names=[\n        \"Chromosome\",\n        \"Start\",\n        \"End\",\n        \"ID1\",\n        \"ID2\",\n        \"Annotation\"\n    ]\n)\n\npromoter_pr = pr.PyRanges(promoter)\nenhancer_pr = pr.PyRanges(enhancer)\n\npromoter_gene = promoter_pr.join(\n    gene_tss_pr\n).df\n\npromoter_gene = promoter_gene[\n    [\n        \"Chromosome\",\n        \"Start\",\n        \"End\",\n        \"gene_id\"\n    ]\n].drop_duplicates()\n\nenhancer_gene = enhancer_pr.nearest(\n    gene_tss_pr,\n    strandedness=False\n).df\n\nenhancer_gene = enhancer_gene[\n    [\n        \"Chromosome\",\n        \"Start\",\n        \"End\",\n        \"gene_id\"\n    ]\n].drop_duplicates()\n\ninput_cells = train_multi_inputs.index.astype(str)\n\nif \"cell_id\" in train_multi_inputs.columns:\n    input_cells = train_multi_inputs[\"cell_id\"].astype(str).values\n    peak_cols = train_multi_inputs.columns[1:]\nelse:\n    peak_cols = train_multi_inputs.columns\n\npeak_cols = pd.Index(peak_cols.astype(str))\n\npeak_info = pd.DataFrame({\n    \"peak\": peak_cols\n})\n\npeak_info[[\"Chromosome\", \"position\"]] = peak_info[\n    \"peak\"\n].str.split(\":\", n=1, expand=True)\n\npeak_info[[\"Start\", \"End\"]] = peak_info[\n    \"position\"\n].str.split(\"-\", n=1, expand=True)\n\npeak_info[\"Start\"] = peak_info[\"Start\"].astype(np.int64)\npeak_info[\"End\"] = peak_info[\"End\"].astype(np.int64)\n\npeak_pr = pr.PyRanges(\n    peak_info[\n        [\n            \"Chromosome\",\n            \"Start\",\n            \"End\",\n            \"peak\"\n        ]\n    ]\n)\n\npeak_promoter = peak_pr.join(\n    pr.PyRanges(promoter_gene)\n).df\n\npeak_promoter = peak_promoter[\n    [\n        \"peak\",\n        \"gene_id\"\n    ]\n].drop_duplicates()\n\npeak_enhancer = peak_pr.join(\n    pr.PyRanges(enhancer_gene)\n).df\n\npeak_enhancer = peak_enhancer[\n    [\n        \"peak\",\n        \"gene_id\"\n    ]\n].drop_duplicates()\n\npromoter_peak_set = set(\n    peak_promoter[\"peak\"]\n)\n\npeak_enhancer = peak_enhancer[\n    ~peak_enhancer[\"peak\"].isin(\n        promoter_peak_set\n    )\n].copy()\n\npromoter_map = peak_promoter.copy()\nenhancer_map = peak_enhancer.copy()\n\npromoter_map[\"gene_symbol\"] = promoter_map[\n    \"gene_id\"\n].map(ensg_to_symbol)\n\nenhancer_map[\"gene_symbol\"] = enhancer_map[\n    \"gene_id\"\n].map(ensg_to_symbol)\n\npromoter_map = promoter_map[\n    promoter_map[\"gene_symbol\"].isin(gene_order)\n].copy()\n\nenhancer_map = enhancer_map[\n    enhancer_map[\"gene_symbol\"].isin(gene_order)\n].copy()\n\ngene_to_idx = {\n    gene: i\n    for i, gene in enumerate(gene_order)\n}\n\npeak_to_idx = {\n    peak: i\n    for i, peak in enumerate(peak_cols)\n}\n\npromoter_map[\"gene_idx\"] = promoter_map[\n    \"gene_symbol\"\n].map(gene_to_idx)\n\npromoter_map[\"peak_idx\"] = promoter_map[\n    \"peak\"\n].map(peak_to_idx)\n\nenhancer_map[\"gene_idx\"] = enhancer_map[\n    \"gene_symbol\"\n].map(gene_to_idx)\n\nenhancer_map[\"peak_idx\"] = enhancer_map[\n    \"peak\"\n].map(peak_to_idx)\n\npromoter_map = promoter_map.dropna(\n    subset=[\"gene_idx\", \"peak_idx\"]\n)\n\nenhancer_map = enhancer_map.dropna(\n    subset=[\"gene_idx\", \"peak_idx\"]\n)\n\npromoter_map[\"gene_idx\"] = promoter_map[\n    \"gene_idx\"\n].astype(np.int64)\n\npromoter_map[\"peak_idx\"] = promoter_map[\n    \"peak_idx\"\n].astype(np.int64)\n\nenhancer_map[\"gene_idx\"] = enhancer_map[\n    \"gene_idx\"\n].astype(np.int64)\n\nenhancer_map[\"peak_idx\"] = enhancer_map[\n    \"peak_idx\"\n].astype(np.int64)\n\nn_genes = len(gene_order)\nn_peaks = len(peak_cols)\n\nP = sp.csr_matrix(\n    (\n        np.ones(len(promoter_map), dtype=np.float32),\n        (\n            promoter_map[\"peak_idx\"].values,\n            promoter_map[\"gene_idx\"].values\n        )\n    ),\n    shape=(n_peaks, n_genes)\n)\n\nE = sp.csr_matrix(\n    (\n        np.ones(len(enhancer_map), dtype=np.float32),\n        (\n            enhancer_map[\"peak_idx\"].values,\n            enhancer_map[\"gene_idx\"].values\n        )\n    ),\n    shape=(n_peaks, n_genes)\n)\n\ninput_cell_to_row = {\n    cell: i\n    for i, cell in enumerate(input_cells)\n}\n\ntarget_cells = target.index.astype(str)\n\ncommon_cells = target_cells[\n    target_cells.isin(input_cell_to_row)\n]\n\ntarget = target.loc[common_cells]\n\ninput_rows = np.array([\n    input_cell_to_row[cell]\n    for cell in common_cells\n])\n\nn_cells = len(common_cells)\n\npromoter_matrix = np.zeros(\n    (n_cells, n_genes),\n    dtype=np.float32\n)\n\nenhancer_matrix = np.zeros(\n    (n_cells, n_genes),\n    dtype=np.float32\n)\n\nchunk_size = 256\n\nfor start in range(0, n_cells, chunk_size):\n\n    end = min(\n        start + chunk_size,\n        n_cells\n    )\n\n    rows = input_rows[start:end]\n\n    if \"cell_id\" in train_multi_inputs.columns:\n        X = train_multi_inputs.iloc[\n            rows, 1:\n        ].to_numpy(\n            dtype=np.float32\n        )\n    else:\n        X = train_multi_inputs.iloc[\n            rows\n        ].to_numpy(\n            dtype=np.float32\n        )\n\n    promoter_matrix[start:end] = (\n        X @ P\n    )\n\n    enhancer_matrix[start:end] = (\n        X @ E\n    )\n\ntarget = target.loc[\n    common_cells,\n    gene_order\n]\n\nrna_matrix = target.to_numpy(\n    dtype=np.float32\n)\n\nadata_multi = ad.AnnData(\n    X=sp.csr_matrix(rna_matrix)\n)\n\nadata_multi.obs_names = common_cells\nadata_multi.var_names = gene_order\n\nadata_multi.layers[\"rna\"] = sp.csr_matrix(\n    rna_matrix\n)\n\nadata_multi.layers[\"A_promoter\"] = sp.csr_matrix(\n    promoter_matrix\n)\n\nadata_multi.layers[\"B_enhancer\"] = sp.csr_matrix(\n    enhancer_matrix\n)"},{"cell_type":"code","execution_count":null,"id":"52553e4d-8b08-40b4-8152-5cce17c83aa8","metadata":{},"outputs":[],"source":"print(\"adata shape:\", adata_multi.shape)\n\nprint(\n    \"RNA:\",\n    adata_multi.layers[\"rna\"].shape\n)\n\nprint(\n    \"A_promoter:\",\n    adata_multi.layers[\"A_promoter\"].shape\n)\n\nprint(\n    \"B_enhancer:\",\n    adata_multi.layers[\"B_enhancer\"].shape\n)\n\nprint(\n    \"cell 정합:\",\n    np.array_equal(\n        adata_multi.obs_names,\n        target.index\n    )\n)\n\nprint(\n    \"gene 정합:\",\n    np.array_equal(\n        adata_multi.var_names,\n        target.columns\n    )\n)\n\nprint(\n    \"RNA-Promoter 정합:\",\n    adata_multi.layers[\"rna\"].shape\n    == adata_multi.layers[\"A_promoter\"].shape\n)\n\nprint(\n    \"RNA-Enhancer 정합:\",\n    adata_multi.layers[\"rna\"].shape\n    == adata_multi.layers[\"B_enhancer\"].shape\n)\n\nprint(\n    \"Promoter-Enhancer 정합:\",\n    adata_multi.layers[\"A_promoter\"].shape\n    == adata_multi.layers[\"B_enhancer\"].shape\n)"},{"cell_type":"code","execution_count":null,"id":"9c1b69d3-78ca-4d48-8b49-7f1aa7af2224","metadata":{},"outputs":[],"source":"metadata = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/metadata.csv\"\n)\n\nmetadata[\"cell_id\"] = metadata[\"cell_id\"].astype(str)\n\nmetadata = metadata.set_index(\"cell_id\")\n\nmetadata = metadata.loc[adata_multi.obs_names]\n\nadata_multi.obs = metadata.copy()"},{"cell_type":"code","execution_count":null,"id":"fb473737-eb98-4475-aa68-0db1e2f17bde","metadata":{},"outputs":[],"source":"output_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/train/train_multi.h5ad\"\nadata_multi.write_h5ad(output_path)"},{"cell_type":"markdown","id":"33514366-04df-448f-a8ba-fd7eda36482c","metadata":{"jp-MarkdownHeadingCollapsed":true},"source":"# Data merge : Multiome(scRNA + CITEseq)\n"},{"cell_type":"code","execution_count":null,"id":"d9c7d5bd-6f86-4b1d-8c9e-606c697d8f93","metadata":{},"outputs":[],"source":"import pandas as pd\n\ninput_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/train/train_cite_inputs.h5\"\ntarget_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/train/train_cite_targets.h5\"\n\ntrain_cite_inputs = pd.read_hdf(\n    input_path,\n    key=\"train_cite_inputs\"\n)\n\ntrain_cite_targets = pd.read_hdf(\n    target_path,\n    key=\"train_cite_targets\"\n)\n\nprint(train_cite_inputs.shape)\nprint(train_cite_targets.shape)"},{"cell_type":"code","execution_count":null,"id":"47c022bd-3fc7-420d-8446-3db0af079227","metadata":{},"outputs":[],"source":"train_cite_inputs.columns = (\n    train_cite_inputs.columns\n    .str.split(\"_\", n=1)\n    .str[1]\n)"},{"cell_type":"code","execution_count":null,"id":"6b7e39a5-a43b-493b-909d-0a31e34e61d1","metadata":{},"outputs":[],"source":"import mudata as md\nimport anndata as ad\nimport scipy.sparse as sp\n\ntrain_cite_inputs.index = train_cite_inputs.index.astype(str)\ntrain_cite_targets.index = train_cite_targets.index.astype(str)\n\ncommon_cells = train_cite_inputs.index.intersection(\n    train_cite_targets.index\n)\n\nrna = train_cite_inputs.loc[common_cells].copy()\nprotein = train_cite_targets.loc[common_cells].copy()\n\nrna.index = common_cells\nprotein.index = common_cells\n\nadata_rna = ad.AnnData(\n    X=sp.csr_matrix(\n        rna.to_numpy(dtype=\"float32\")\n    )\n)\n\nadata_rna.obs_names = common_cells\nadata_rna.var_names = rna.columns.astype(str)\n\nadata_protein = ad.AnnData(\n    X=sp.csr_matrix(\n        protein.to_numpy(dtype=\"float32\")\n    )\n)\n\nadata_protein.obs_names = common_cells\nadata_protein.var_names = protein.columns.astype(str)\n\nmdata_cite = md.MuData({\n    \"rna\": adata_rna,\n    \"protein\": adata_protein\n})\n\nprint(mdata_cite)\nprint(mdata_cite[\"rna\"].shape)\nprint(mdata_cite[\"protein\"].shape)"},{"cell_type":"code","execution_count":null,"id":"f1ff3f9b-7301-4c18-8464-7e89f448fa48","metadata":{},"outputs":[],"source":"metadata = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/metadata.csv\"\n)\nmdata_cite.obs = metadata.set_index(\"cell_id\").loc[mdata_cite[\"rna\"].obs_names].copy()"},{"cell_type":"code","execution_count":null,"id":"6c296eaf-4209-4c3d-ba0e-589e4b35e5df","metadata":{},"outputs":[],"source":"output_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/train/train_cite.h5ad\"\n\nmdata_cite.write_h5mu(output_path)"},{"cell_type":"markdown","id":"5def0d13-b2b4-47c3-bdba-fa5e1f3fd512","metadata":{"jp-MarkdownHeadingCollapsed":true},"source":"# Test data"},{"cell_type":"code","execution_count":null,"id":"9448501a-342b-4741-ae56-f31d75e7282c","metadata":{},"outputs":[],"source":"import pandas as pd\n\ntest_cite_day2 = pd.read_hdf(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/test/test_cite_inputs_day_2_donor_27678.h5\",\n    key=\"test_cite_inputs_day_2_donor_27678\"\n)\n\ntest_cite = pd.read_hdf(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/test/test_cite_inputs.h5\",\n    key=\"test_cite_inputs\"\n)\n\ntest_multi = pd.read_hdf(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/test/test_multi_inputs.h5\",\n    key=\"test_multi_inputs\"\n)\n\nprint(test_cite_day2.shape)\nprint(test_cite.shape)\nprint(test_multi.shape)"},{"cell_type":"code","execution_count":null,"id":"3c986c1f-91e9-4e1b-9aac-22b41855d680","metadata":{},"outputs":[],"source":"test_cite_day2.columns = (\n    test_cite_day2.columns\n    .str.split(\"_\", n=1)\n    .str[1])\n\ntest_cite.columns = (\n    test_cite.columns\n    .str.split(\"_\", n=1)\n    .str[1])"},{"cell_type":"code","execution_count":null,"id":"74ee1508-a045-4726-b47a-45bae8ded226","metadata":{},"outputs":[],"source":"test_cite_day2 = test_cite_day2.loc[:, ~test_cite_day2.columns.duplicated()]\ntest_cite = test_cite.loc[:, ~test_cite.columns.duplicated()]"},{"cell_type":"code","execution_count":null,"id":"df03d6cf-0848-4339-ba63-b299937d9871","metadata":{},"outputs":[],"source":"common_cols = test_cite_day2.columns.intersection(test_cite.columns)\n\ntest_cite_combined = pd.concat(\n    [\n        test_cite_day2[common_cols],\n        test_cite[common_cols]\n    ],\n    axis=0,\n    ignore_index=True\n)"},{"cell_type":"code","execution_count":null,"id":"1fa5221d-3448-49e4-a98b-03b8f6a143af","metadata":{},"outputs":[],"source":"import pandas as pd\n\ncommon_cols = test_cite_day2.columns.intersection(test_cite.columns)\n\ntest_cite_combined = pd.concat(\n    [\n        test_cite_day2[common_cols],\n        test_cite[common_cols]\n    ],\n    axis=0,\n    join=\"inner\"\n)\n\nprint(\"test_cite_day2:\", test_cite_day2.shape)\nprint(\"test_cite:\", test_cite.shape)\nprint(\"공통 열 수:\", len(common_cols))\nprint(\"매칭되지 않은 test_cite_day2 열:\", len(test_cite_day2.columns.difference(common_cols)))\nprint(\"매칭되지 않은 test_cite 열:\", len(test_cite.columns.difference(common_cols)))\nprint(\"합친 데이터:\", test_cite_combined.shape)"},{"cell_type":"code","execution_count":null,"id":"bed6871b-bb2c-4919-ae78-b59e4ed97b17","metadata":{},"outputs":[],"source":"import pandas as pd\n\nmetadata = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/metadata.csv\"\n)\n\nmetadata_cite_day2 = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/metadata_cite_day_2_donor_27678.csv\"\n)\n\nmetadata[\"cell_id\"] = metadata[\"cell_id\"].astype(str)\nmetadata_cite_day2[\"cell_id\"] = metadata_cite_day2[\"cell_id\"].astype(str)\ntest_cite_combined.index = test_cite_combined.index.astype(str)\n\nmatched_metadata = test_cite_combined.index.intersection(\n    metadata[\"cell_id\"]\n)\n\nmatched_day2 = test_cite_combined.index.intersection(\n    metadata_cite_day2[\"cell_id\"]\n)\n\nprint(\"test_cite_combined cell 수:\", len(test_cite_combined))\nprint(\"metadata.csv cell 수:\", len(metadata))\nprint(\"metadata.csv 매칭되는 cell 수:\", len(matched_metadata))\nprint()\nprint(\"metadata_cite_day_2_donor_27678.csv cell 수:\", len(metadata_cite_day2))\nprint(\"metadata_cite_day_2_donor_27678.csv 매칭되는 cell 수:\", len(matched_day2))"},{"cell_type":"code","execution_count":null,"id":"43571573-65fc-4d65-9343-612fbad1a3f9","metadata":{},"outputs":[],"source":"import pandas as pd\nimport anndata as ad\nimport scipy.sparse as sp\n\nmetadata = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/metadata.csv\"\n)\n\nmetadata_cite_day2 = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/metadata_cite_day_2_donor_27678.csv\"\n)\n\nmetadata[\"cell_id\"] = metadata[\"cell_id\"].astype(str)\nmetadata_cite_day2[\"cell_id\"] = metadata_cite_day2[\"cell_id\"].astype(str)\n\ntest_cite_combined.index = test_cite_combined.index.astype(str)\n\nmetadata_all = pd.concat(\n    [\n        metadata.set_index(\"cell_id\"),\n        metadata_cite_day2.set_index(\"cell_id\")\n    ],\n    axis=0\n)\n\nmetadata_all = metadata_all.loc[\n    ~metadata_all.index.duplicated(keep=\"first\")\n]\n\nmatched_obs = metadata_all.loc[test_cite_combined.index]\n\nadata_cite_test = ad.AnnData(\n    X=sp.csr_matrix(test_cite_combined.to_numpy(dtype=\"float32\"))\n)\n\nadata_cite_test.obs_names = test_cite_combined.index\nadata_cite_test.var_names = test_cite_combined.columns.astype(str)\nadata_cite_test.obs = matched_obs.copy()\n\nprint(adata_cite_test)\nprint(adata_cite_test.obs.head())"},{"cell_type":"code","execution_count":null,"id":"a8ed1f66-b56e-4189-af7e-2c57ed1af706","metadata":{},"outputs":[],"source":"output_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/test/test_rna.h5ad\"\n\nadata_cite_test.write_h5ad(output_path)"},{"cell_type":"code","execution_count":null,"id":"5adcebc4-e219-412c-bce9-a3cbfb179857","metadata":{},"outputs":[],"source":"#### ATAC split ####"},{"cell_type":"code","execution_count":null,"id":"0cc20bc6-d660-43ff-8988-61d3f9edfaac","metadata":{},"outputs":[],"source":"import pandas as pd\nimport numpy as np\nimport scipy.sparse as sp\nimport anndata as ad\nimport pyranges as pr\n\nbase = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/baseline\"\n\npromoter_path = f\"{base}/promoter.bed\"\nenhancer_path = f\"{base}/enhancer.bed\"\ngtf_path = f\"{base}/gencode.v44.annotation.gtf.gz\"\n\ngtf = pd.read_csv(\n    gtf_path,\n    sep=\"\\t\",\n    comment=\"#\",\n    header=None,\n    names=[\n        \"Chromosome\", \"Source\", \"Feature\", \"Start\", \"End\",\n        \"Score\", \"Strand\", \"Frame\", \"Attribute\"\n    ]\n)\n\ngtf_gene = gtf[gtf[\"Feature\"] == \"gene\"].copy()\n\ngtf_gene[\"gene_id\"] = gtf_gene[\"Attribute\"].str.extract(\n    r'gene_id \"([^\"]+)\"'\n)[0]\n\ngtf_gene[\"gene_name\"] = gtf_gene[\"Attribute\"].str.extract(\n    r'gene_name \"([^\"]+)\"'\n)[0]\n\ngtf_gene[\"gene_id\"] = gtf_gene[\"gene_id\"].str.replace(\n    r\"\\.\\d+$\", \"\", regex=True\n)\n\nensg_to_symbol = (\n    gtf_gene[\n        [\"gene_id\", \"gene_name\"]\n    ]\n    .dropna()\n    .drop_duplicates(\"gene_id\")\n    .set_index(\"gene_id\")[\"gene_name\"]\n    .to_dict()\n)\n\ngene_order = pd.Index(\n    gtf_gene[\"gene_name\"].dropna().drop_duplicates().astype(str)\n)\n\ngene_tss = gtf_gene[\n    [\n        \"Chromosome\",\n        \"Start\",\n        \"End\",\n        \"Strand\",\n        \"gene_id\",\n        \"gene_name\"\n    ]\n].copy()\n\ngene_tss[\"Start\"] = np.where(\n    gene_tss[\"Strand\"] == \"+\",\n    gene_tss[\"Start\"] - 1,\n    gene_tss[\"End\"] - 1\n)\n\ngene_tss[\"End\"] = gene_tss[\"Start\"] + 1\n\ngene_tss_pr = pr.PyRanges(gene_tss)\n\npromoter = pd.read_csv(\n    promoter_path,\n    sep=\"\\t\",\n    header=None,\n    names=[\n        \"Chromosome\",\n        \"Start\",\n        \"End\",\n        \"ID1\",\n        \"ID2\",\n        \"Annotation\"\n    ]\n)\n\nenhancer = pd.read_csv(\n    enhancer_path,\n    sep=\"\\t\",\n    header=None,\n    names=[\n        \"Chromosome\",\n        \"Start\",\n        \"End\",\n        \"ID1\",\n        \"ID2\",\n        \"Annotation\"\n    ]\n)\n\npromoter_gene = pr.PyRanges(promoter).join(\n    gene_tss_pr\n).df[\n    [\"Chromosome\", \"Start\", \"End\", \"gene_id\"]\n].drop_duplicates()\n\nenhancer_gene = pr.PyRanges(enhancer).nearest(\n    gene_tss_pr,\n    strandedness=False\n).df[\n    [\"Chromosome\", \"Start\", \"End\", \"gene_id\"]\n].drop_duplicates()\n\nif \"cell_id\" in test_multi.columns:\n    cell_ids = test_multi[\"cell_id\"].astype(str).values\n    peak_cols = pd.Index(test_multi.columns[1:].astype(str))\n    X = test_multi.iloc[:, 1:]\nelse:\n    cell_ids = test_multi.index.astype(str).values\n    peak_cols = pd.Index(test_multi.columns.astype(str))\n    X = test_multi\n\npeak_info = pd.DataFrame({\n    \"peak\": peak_cols\n})\n\npeak_info[[\"Chromosome\", \"position\"]] = peak_info[\n    \"peak\"\n].str.split(\":\", n=1, expand=True)\n\npeak_info[[\"Start\", \"End\"]] = peak_info[\n    \"position\"\n].str.split(\"-\", n=1, expand=True)\n\npeak_info[\"Start\"] = peak_info[\"Start\"].astype(np.int64)\npeak_info[\"End\"] = peak_info[\"End\"].astype(np.int64)\n\npeak_pr = pr.PyRanges(\n    peak_info[\n        [\"Chromosome\", \"Start\", \"End\", \"peak\"]\n    ]\n)\n\npeak_promoter = peak_pr.join(\n    pr.PyRanges(promoter_gene)\n).df[\n    [\"peak\", \"gene_id\"]\n].drop_duplicates()\n\npeak_enhancer = peak_pr.join(\n    pr.PyRanges(enhancer_gene)\n).df[\n    [\"peak\", \"gene_id\"]\n].drop_duplicates()\n\npromoter_peak_set = set(\n    peak_promoter[\"peak\"]\n)\n\npeak_enhancer = peak_enhancer[\n    ~peak_enhancer[\"peak\"].isin(promoter_peak_set)\n].copy()\n\npromoter_map = peak_promoter.copy()\nenhancer_map = peak_enhancer.copy()\n\npromoter_map[\"gene_symbol\"] = promoter_map[\"gene_id\"].map(\n    ensg_to_symbol\n)\n\nenhancer_map[\"gene_symbol\"] = enhancer_map[\"gene_id\"].map(\n    ensg_to_symbol\n)\n\npromoter_map = promoter_map[\n    promoter_map[\"gene_symbol\"].isin(gene_order)\n].copy()\n\nenhancer_map = enhancer_map[\n    enhancer_map[\"gene_symbol\"].isin(gene_order)\n].copy()\n\ngene_to_idx = {\n    gene: i\n    for i, gene in enumerate(gene_order)\n}\n\npeak_to_idx = {\n    peak: i\n    for i, peak in enumerate(peak_cols)\n}\n\npromoter_map[\"gene_idx\"] = promoter_map[\"gene_symbol\"].map(\n    gene_to_idx\n)\n\npromoter_map[\"peak_idx\"] = promoter_map[\"peak\"].map(\n    peak_to_idx\n)\n\nenhancer_map[\"gene_idx\"] = enhancer_map[\"gene_symbol\"].map(\n    gene_to_idx\n)\n\nenhancer_map[\"peak_idx\"] = enhancer_map[\"peak\"].map(\n    peak_to_idx\n)\n\npromoter_map = promoter_map.dropna(\n    subset=[\"gene_idx\", \"peak_idx\"]\n)\n\nenhancer_map = enhancer_map.dropna(\n    subset=[\"gene_idx\", \"peak_idx\"]\n)\n\npromoter_map[\"gene_idx\"] = promoter_map[\"gene_idx\"].astype(np.int64)\npromoter_map[\"peak_idx\"] = promoter_map[\"peak_idx\"].astype(np.int64)\n\nenhancer_map[\"gene_idx\"] = enhancer_map[\"gene_idx\"].astype(np.int64)\nenhancer_map[\"peak_idx\"] = enhancer_map[\"peak_idx\"].astype(np.int64)\n\nn_cells = len(cell_ids)\nn_genes = len(gene_order)\nn_peaks = len(peak_cols)\n\nP = sp.csr_matrix(\n    (\n        np.ones(len(promoter_map), dtype=np.float32),\n        (\n            promoter_map[\"peak_idx\"].values,\n            promoter_map[\"gene_idx\"].values\n        )\n    ),\n    shape=(n_peaks, n_genes)\n)\n\nE = sp.csr_matrix(\n    (\n        np.ones(len(enhancer_map), dtype=np.float32),\n        (\n            enhancer_map[\"peak_idx\"].values,\n            enhancer_map[\"gene_idx\"].values\n        )\n    ),\n    shape=(n_peaks, n_genes)\n)\n\npromoter_matrix = np.zeros(\n    (n_cells, n_genes),\n    dtype=np.float32\n)\n\nenhancer_matrix = np.zeros(\n    (n_cells, n_genes),\n    dtype=np.float32\n)\n\nchunk_size = 256\n\nfor start in range(0, n_cells, chunk_size):\n    end = min(start + chunk_size, n_cells)\n\n    X_chunk = X.iloc[\n        start:end\n    ].to_numpy(\n        dtype=np.float32\n    )\n\n    promoter_matrix[start:end] = X_chunk @ P\n    enhancer_matrix[start:end] = X_chunk @ E\n\nadata_multi_test = ad.AnnData(\n    X=sp.csr_matrix(promoter_matrix)\n)\n\nadata_multi_test.obs_names = cell_ids\nadata_multi_test.var_names = gene_order\n\nadata_multi_test.layers[\"A_promoter\"] = sp.csr_matrix(\n    promoter_matrix\n)\n\nadata_multi_test.layers[\"B_enhancer\"] = sp.csr_matrix(\n    enhancer_matrix\n)\n\nprint(adata_multi_test)\nprint(adata_multi_test.layers.keys())\nprint(adata_multi_test.layers[\"A_promoter\"].shape)\nprint(adata_multi_test.layers[\"B_enhancer\"].shape)"},{"cell_type":"code","execution_count":null,"id":"6223efd5-d57f-4338-8a6e-2c49902abf2b","metadata":{},"outputs":[],"source":"metadata = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/metadata.csv\"\n)\n\nmetadata[\"cell_id\"] = metadata[\"cell_id\"].astype(str)\nmetadata = metadata.set_index(\"cell_id\")\n\nadata_multi_test.obs = metadata.loc[\n    adata_multi_test.obs_names.astype(str)\n].copy()"},{"cell_type":"code","execution_count":null,"id":"652911b7-79a5-4628-bc6a-40b197d822aa","metadata":{},"outputs":[],"source":"output_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/test/test_atac.h5ad\"\n\nadata_multi_test.write_h5ad(output_path)"},{"cell_type":"code","execution_count":null,"id":"20ae2870-8a25-4a24-aa8e-216c4d3b5562","metadata":{},"outputs":[],"source":"import mudata as md\n\ntrain_cite = md.read_h5mu(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/train/train_cite.h5ad\"\n)\n\nprint(train_cite)"},{"cell_type":"code","execution_count":null,"id":"83f3ea96-1e50-4d05-bf93-88f6fd083f31","metadata":{},"outputs":[],"source":"train_cite.var"},{"cell_type":"code","execution_count":null,"id":"670076c3-dbed-4313-94b3-a8c7b2e4eec6","metadata":{},"outputs":[],"source":""},{"cell_type":"markdown","id":"3a874a5f-d4a8-40c5-ae98-0b86290d8e44","metadata":{"jp-MarkdownHeadingCollapsed":true},"source":"# Data merge"},{"cell_type":"code","execution_count":5,"id":"1469fac3-4ce4-47c3-a3c0-0a37742d5639","metadata":{},"outputs":[],"source":"import mudata as md\nimport anndata as ad\n\ntrain_cite = md.read_h5mu(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/train/train_cite.h5ad\"\n)\n\ntrain_multi = ad.read_h5ad(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/train/train_multi.h5ad\"\n)\n\ntest_rna = ad.read_h5ad(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/test/test_rna.h5ad\"\n)\n\ntest_atac = ad.read_h5ad(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/test/test_atac.h5ad\"\n)\n\ngenes = {\n    \"train_cite\": set(train_cite[\"rna\"].var_names),\n    \"train_multi\": set(train_multi.var_names),\n    \"test_rna\": set(test_rna.var_names),\n    \"test_atac\": set(test_atac.var_names)\n}\n\nnames = list(genes.keys())\n\nfor i in range(len(names)):\n    for j in range(i + 1, len(names)):\n        a, b = names[i], names[j]\n\n        overlap = genes[a] & genes[b]\n        only_a = genes[a] - genes[b]\n        only_b = genes[b] - genes[a]\n\n        print(f\"{a} vs {b}\")\n        print(f\"  {a}: {len(genes[a])}\")\n        print(f\"  {b}: {len(genes[b])}\")\n        print(f\"  겹치는 유전자: {len(overlap)}\")\n        print(f\"  {a}에만 존재: {len(only_a)}\")\n        print(f\"  {b}에만 존재: {len(only_b)}\")\n        print()\n\ncommon_all = set.intersection(*genes.values())\n\nprint(\"=\" * 50)\nprint(f\"4개 데이터셋 공통 유전자: {len(common_all)}\")"},{"cell_type":"code","execution_count":6,"id":"55c96b94-cf1f-423e-ad44-bfe45dd6c370","metadata":{},"outputs":[],"source":"A = train_multi.layers[\"A_promoter\"].toarray()\nB = train_multi.layers[\"B_enhancer\"].toarray()\n\ntrain_multi.layers[\"peak_sum\"] = (1 + A) * (1 + B)\ntrain_multi"},{"cell_type":"code","execution_count":7,"id":"a2fe4096-6e26-4c01-97fd-ac09ca8317a0","metadata":{},"outputs":[],"source":"A = test_atac.layers[\"A_promoter\"].toarray()\nB = test_atac.layers[\"B_enhancer\"].toarray()\n\ntest_atac.layers[\"peak_sum\"] = (1 + A) * (1 + B)\ntest_atac"},{"cell_type":"code","execution_count":8,"id":"bbf57ffb-7f49-46d3-83a4-4291845d7192","metadata":{},"outputs":[],"source":"test_atac = test_atac[:, train_multi.var_names].copy()"},{"cell_type":"code","execution_count":9,"id":"4c4d292f-6f7f-4b63-a2e1-f7589c6053ce","metadata":{},"outputs":[],"source":"print(\"=== train_multi ===\")\nprint(\"Shape:\", train_multi.shape)\nprint(\"Duplicated genes:\", train_multi.var_names.duplicated().sum())\nprint(\n    \"Unique duplicated genes:\",\n    train_multi.var_names[train_multi.var_names.duplicated(keep=False)].unique().size\n)\n\nprint(\"\\n=== train_cite RNA ===\")\nprint(\"Shape:\", train_cite[\"rna\"].shape)\nprint(\"Duplicated genes:\", train_cite[\"rna\"].var_names.duplicated().sum())\nprint(\n    \"Unique duplicated genes:\",\n    train_cite[\"rna\"].var_names[\n        train_cite[\"rna\"].var_names.duplicated(keep=False)\n    ].unique().size\n)\n\nprint(\"\\n=== test_atac ===\")\nprint(\"Shape:\", test_atac.shape)\nprint(\"Duplicated genes:\", test_atac.var_names.duplicated().sum())\nprint(\n    \"Unique duplicated genes:\",\n    test_atac.var_names[\n        test_atac.var_names.duplicated(keep=False)\n    ].unique().size\n)\n\nprint(\"\\n=== test_rna ===\")\nprint(\"Shape:\", test_rna.shape)\nprint(\"Duplicated genes:\", test_rna.var_names.duplicated().sum())\nprint(\n    \"Unique duplicated genes:\",\n    test_rna.var_names[\n        test_rna.var_names.duplicated(keep=False)\n    ].unique().size\n)"},{"cell_type":"code","execution_count":10,"id":"2dc7a516-0085-4edc-b4a2-23ff3e3a9d78","metadata":{},"outputs":[],"source":"import numpy as np\n\n# CITE RNA: 중복 gene은 첫 번째 occurrence만 유지\ntrain_cite_rna = train_cite[\"rna\"]\n\ncite_keep = ~train_cite_rna.var_names.duplicated(keep=\"first\")\n\ntrain_cite_rna_unique = train_cite_rna[:, cite_keep].copy()\n\n# Multiome: 중복 없음\ntrain_multi_unique = train_multi\n\n# --------------------------------------------------\n# 최종 gene 개수 확인\n# --------------------------------------------------\n\nprint(\"=== Gene count ===\")\nprint(\"train_cite RNA:\", train_cite_rna_unique.n_vars)\nprint(\"test_rna      :\", test_rna.n_vars)\n\nprint(\"train_multi   :\", train_multi_unique.n_vars)\nprint(\"test_atac     :\", test_atac.n_vars)\n\n\n# --------------------------------------------------\n# Gene 순서 확인\n# --------------------------------------------------\n\ncite_order_match = np.array_equal(\n    train_cite_rna_unique.var_names.to_numpy(),\n    test_rna.var_names.to_numpy()\n)\n\nmulti_order_match = np.array_equal(\n    train_multi_unique.var_names.to_numpy(),\n    test_atac.var_names.to_numpy()\n)\n\nprint(\"\\n=== Gene count & order check ===\")\nprint(\n    \"train_cite RNA ↔ test_rna:\",\n    cite_order_match\n)\n\nprint(\n    \"train_multi ↔ test_atac:\",\n    multi_order_match\n)\n\n\n# --------------------------------------------------\n# 개수와 순서가 모두 일치하는지 최종 확인\n# --------------------------------------------------\n\nassert train_cite_rna_unique.n_vars == test_rna.n_vars\nassert train_multi_unique.n_vars == test_atac.n_vars\n\nassert cite_order_match\nassert multi_order_match\n\nprint(\"\\nAll gene features match.\")"},{"cell_type":"code","execution_count":null,"id":"b0827e56-94c8-4a36-a471-3af88adfb025","metadata":{},"outputs":[],"source":""},{"cell_type":"markdown","id":"54d89029-6802-420d-a134-3a7cfbcbb51a","metadata":{},"source":"# Multimodal diffusion modeling"},{"cell_type":"code","execution_count":11,"id":"ad36d510-274e-46be-a629-bb47d914a4ea","metadata":{},"outputs":[],"source":"import math\nimport itertools\nimport numpy as np\nimport torch\nimport os\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nprint(\"=\" * 60)\nprint(\"PyTorch version:\", torch.__version__)\nprint(\"CUDA available:\", torch.cuda.is_available())\nprint(\"CUDA version:\", torch.version.cuda)\nprint(\"GPU count:\", torch.cuda.device_count())\n\nif torch.cuda.is_available():\n    print(\"GPU:\", torch.cuda.get_device_name(0))\n    print(\"Current GPU:\", torch.cuda.current_device())\n    print(\n        \"GPU memory:\",\n        f\"{torch.cuda.get_device_properties(0).total_memory / 1024**3:.2f} GB\"\n    )\n    device = torch.device(\"cuda\")\nelse:\n    print(\"GPU: None\")\n    device = torch.device(\"cpu\")\n\nprint(\"Using device:\", device)\nprint(\"=\" * 60)"},{"cell_type":"code","execution_count":12,"id":"8ca76696-1f80-4787-9710-3c55e05239d0","metadata":{},"outputs":[],"source":"import os\nimport math\nimport itertools\nimport numpy as np\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\ntrain_cite_rna = train_cite[\"rna\"]\ncite_keep = ~train_cite_rna.var_names.duplicated(keep=\"first\")\ntrain_cite_rna_unique = train_cite_rna[:, cite_keep].copy()\ntrain_multi_unique = train_multi\n\nprint(\"=== Gene count ===\")\nprint(\"train_cite RNA original:\", train_cite_rna.n_vars)\nprint(\"train_cite RNA unique  :\", train_cite_rna_unique.n_vars)\nprint(\"test_rna               :\", test_rna.n_vars)\nprint(\"train_multi            :\", train_multi_unique.n_vars)\nprint(\"test_atac              :\", test_atac.n_vars)\n\ncite_order_match = np.array_equal(\n    train_cite_rna_unique.var_names.to_numpy(),\n    test_rna.var_names.to_numpy()\n)\n\nmulti_order_match = np.array_equal(\n    train_multi_unique.var_names.to_numpy(),\n    test_atac.var_names.to_numpy()\n)\n\nprint(\"\\n=== Gene count & order check ===\")\nprint(\"train_cite RNA unique ↔ test_rna:\", cite_order_match)\nprint(\"train_multi ↔ test_atac:\", multi_order_match)\n\nassert train_cite_rna_unique.n_vars == test_rna.n_vars\nassert train_multi_unique.n_vars == test_atac.n_vars\nassert cite_order_match\nassert multi_order_match\n\nprint(\"\\nAll gene features match.\")\n\ndevice = torch.device(\n    \"cuda\" if torch.cuda.is_available() else \"cpu\"\n)\n\nmulti_gene_dim = train_multi_unique.layers[\"rna\"].shape[1]\ncite_gene_dim = train_cite_rna_unique.shape[1]\nprotein_dim = train_cite[\"protein\"].shape[1]\n\nlatent_dim = 128\nprotein_latent_dim = 64\nhidden_dim = 512\ntime_dim = 64\nnum_diffusion_steps = 200\n\nbatch_size = 128\nlr = 1e-4\nepochs = 50\npatience = 10\n\n\nclass SinusoidalTimeEmbedding(nn.Module):\n\n    def __init__(self, dim):\n        super().__init__()\n        self.dim = dim\n\n    def forward(self, t):\n        half = self.dim // 2\n        emb = math.log(10000) / (half - 1)\n\n        emb = torch.exp(\n            torch.arange(\n                half,\n                device=t.device\n            ) * -emb\n        )\n\n        emb = (\n            t.float().unsqueeze(1)\n            * emb.unsqueeze(0)\n        )\n\n        return torch.cat(\n            [emb.sin(), emb.cos()],\n            dim=1\n        )\n\n\nclass Encoder(nn.Module):\n\n    def __init__(self, input_dim, latent_dim):\n        super().__init__()\n\n        self.net = nn.Sequential(\n            nn.Linear(input_dim, hidden_dim),\n            nn.GELU(),\n            nn.LayerNorm(hidden_dim),\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.GELU(),\n            nn.Linear(hidden_dim, latent_dim)\n        )\n\n    def forward(self, x):\n        return self.net(x)\n\n\nclass Decoder(nn.Module):\n\n    def __init__(self, latent_dim, output_dim):\n        super().__init__()\n\n        self.net = nn.Sequential(\n            nn.Linear(latent_dim, hidden_dim),\n            nn.GELU(),\n            nn.LayerNorm(hidden_dim),\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.GELU(),\n            nn.Linear(hidden_dim, output_dim)\n        )\n\n    def forward(self, z):\n        return self.net(z)\n\n\nclass ConditionalDenoiser(nn.Module):\n\n    def __init__(\n        self,\n        latent_dim,\n        time_dim\n    ):\n        super().__init__()\n\n        self.time_embedding = nn.Sequential(\n            SinusoidalTimeEmbedding(time_dim),\n            nn.Linear(time_dim, time_dim),\n            nn.GELU()\n        )\n\n        input_dim = (\n            latent_dim +\n            latent_dim +\n            time_dim\n        )\n\n        self.net = nn.Sequential(\n            nn.Linear(input_dim, hidden_dim),\n            nn.GELU(),\n            nn.LayerNorm(hidden_dim),\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.GELU(),\n            nn.Linear(hidden_dim, latent_dim)\n        )\n\n    def forward(\n        self,\n        x_t,\n        condition,\n        t\n    ):\n        t_emb = self.time_embedding(t)\n\n        x = torch.cat(\n            [\n                x_t,\n                condition,\n                t_emb\n            ],\n            dim=1\n        )\n\n        return self.net(x)\n\n\nclass ProteinConditionalDenoiser(nn.Module):\n\n    def __init__(\n        self,\n        protein_latent_dim,\n        rna_latent_dim,\n        time_dim\n    ):\n        super().__init__()\n\n        self.time_embedding = nn.Sequential(\n            SinusoidalTimeEmbedding(time_dim),\n            nn.Linear(time_dim, time_dim),\n            nn.GELU()\n        )\n\n        input_dim = (\n            protein_latent_dim +\n            rna_latent_dim +\n            time_dim\n        )\n\n        self.net = nn.Sequential(\n            nn.Linear(input_dim, hidden_dim),\n            nn.GELU(),\n            nn.LayerNorm(hidden_dim),\n            nn.Linear(hidden_dim, hidden_dim),\n            nn.GELU(),\n            nn.Linear(hidden_dim, protein_latent_dim)\n        )\n\n    def forward(\n        self,\n        x_t,\n        rna_condition,\n        t\n    ):\n        t_emb = self.time_embedding(t)\n\n        x = torch.cat(\n            [\n                x_t,\n                rna_condition,\n                t_emb\n            ],\n            dim=1\n        )\n\n        return self.net(x)\n\n\nclass DiffusionSchedule:\n\n    def __init__(\n        self,\n        num_steps,\n        device\n    ):\n        self.num_steps = num_steps\n\n        self.beta = torch.linspace(\n            1e-4,\n            0.02,\n            num_steps,\n            device=device\n        )\n\n        self.alpha = 1.0 - self.beta\n\n        self.alpha_bar = torch.cumprod(\n            self.alpha,\n            dim=0\n        )\n\n    def q_sample(\n        self,\n        x0,\n        t,\n        noise=None\n    ):\n        if noise is None:\n            noise = torch.randn_like(x0)\n\n        alpha_bar_t = (\n            self.alpha_bar[t]\n            .unsqueeze(1)\n        )\n\n        x_t = (\n            torch.sqrt(alpha_bar_t) * x0 +\n            torch.sqrt(1.0 - alpha_bar_t) * noise\n        )\n\n        return x_t, noise\n\n\nclass MultimodalDiffusionModel(nn.Module):\n\n    def __init__(\n        self,\n        multi_gene_dim,\n        cite_gene_dim,\n        protein_dim\n    ):\n        super().__init__()\n\n        self.atac_encoder = Encoder(\n            multi_gene_dim,\n            latent_dim\n        )\n\n        self.multi_rna_encoder = Encoder(\n            multi_gene_dim,\n            latent_dim\n        )\n\n        self.cite_rna_encoder = Encoder(\n            cite_gene_dim,\n            latent_dim\n        )\n\n        self.protein_encoder = Encoder(\n            protein_dim,\n            protein_latent_dim\n        )\n\n        self.multi_rna_decoder = Decoder(\n            latent_dim,\n            multi_gene_dim\n        )\n\n        self.cite_rna_decoder = Decoder(\n            latent_dim,\n            cite_gene_dim\n        )\n\n        self.protein_decoder = Decoder(\n            protein_latent_dim,\n            protein_dim\n        )\n\n        self.atac_to_rna_diffusion = ConditionalDenoiser(\n            latent_dim,\n            time_dim\n        )\n\n        self.rna_to_protein_diffusion = ProteinConditionalDenoiser(\n            protein_latent_dim,\n            latent_dim,\n            time_dim\n        )\n\n    def encode_atac(self, x):\n        return self.atac_encoder(x)\n\n    def encode_multi_rna(self, x):\n        return self.multi_rna_encoder(x)\n\n    def encode_cite_rna(self, x):\n        return self.cite_rna_encoder(x)\n\n    def encode_protein(self, x):\n        return self.protein_encoder(x)\n\n\nclass MultiomeDataset(Dataset):\n\n    def __init__(\n        self,\n        peak,\n        rna\n    ):\n        self.peak = peak\n        self.rna = rna\n\n    def __len__(self):\n        return self.peak.shape[0]\n\n    def __getitem__(self, idx):\n\n        peak = self.peak[idx]\n        rna = self.rna[idx]\n\n        if hasattr(peak, \"toarray\"):\n            peak = peak.toarray().ravel()\n\n        if hasattr(rna, \"toarray\"):\n            rna = rna.toarray().ravel()\n\n        peak = np.asarray(\n            peak,\n            dtype=np.float32\n        )\n\n        rna = np.asarray(\n            rna,\n            dtype=np.float32\n        )\n\n        return (\n            torch.from_numpy(peak),\n            torch.from_numpy(rna)\n        )\n\n\nclass CiteDataset(Dataset):\n\n    def __init__(\n        self,\n        rna,\n        protein\n    ):\n        self.rna = rna\n        self.protein = protein\n\n    def __len__(self):\n        return self.rna.shape[0]\n\n    def __getitem__(self, idx):\n\n        rna = self.rna[idx]\n        protein = self.protein[idx]\n\n        if hasattr(rna, \"toarray\"):\n            rna = rna.toarray().ravel()\n\n        if hasattr(protein, \"toarray\"):\n            protein = protein.toarray().ravel()\n\n        rna = np.asarray(\n            rna,\n            dtype=np.float32\n        )\n\n        protein = np.asarray(\n            protein,\n            dtype=np.float32\n        )\n\n        return (\n            torch.from_numpy(rna),\n            torch.from_numpy(protein)\n        )\n\n\nmulti_peak = train_multi_unique.layers[\"peak_sum\"]\nmulti_rna = train_multi_unique.layers[\"rna\"]\n\ncite_rna = train_cite_rna_unique.X\ncite_protein = train_cite[\"protein\"].X\n\nprint(\"\\n=== Training data ===\")\nprint(\"Multiome peak:\", multi_peak.shape)\nprint(\"Multiome RNA:\", multi_rna.shape)\nprint(\"CITE RNA:\", cite_rna.shape)\nprint(\"CITE protein:\", cite_protein.shape)\n\n\nmulti_dataset = MultiomeDataset(\n    multi_peak,\n    multi_rna\n)\n\ncite_dataset = CiteDataset(\n    cite_rna,\n    cite_protein\n)\n\n\nmulti_loader = DataLoader(\n    multi_dataset,\n    batch_size=batch_size,\n    shuffle=True,\n    num_workers=4,\n    pin_memory=True,\n    drop_last=True,\n    persistent_workers=True\n)\n\ncite_loader = DataLoader(\n    cite_dataset,\n    batch_size=batch_size,\n    shuffle=True,\n    num_workers=4,\n    pin_memory=True,\n    drop_last=True,\n    persistent_workers=True\n)\n\n\nmodel = MultimodalDiffusionModel(\n    multi_gene_dim=multi_gene_dim,\n    cite_gene_dim=cite_gene_dim,\n    protein_dim=protein_dim\n).to(device)\n\n\nschedule = DiffusionSchedule(\n    num_diffusion_steps,\n    device\n)\n\n\noptimizer = torch.optim.AdamW(\n    model.parameters(),\n    lr=lr,\n    weight_decay=1e-5\n)\n\n\ndef multiome_loss(\n    model,\n    peak,\n    rna\n):\n\n    peak = peak.to(\n        device,\n        non_blocking=True\n    )\n\n    rna = rna.to(\n        device,\n        non_blocking=True\n    )\n\n    z_atac = model.encode_atac(peak)\n    z_rna = model.encode_multi_rna(rna)\n\n    t = torch.randint(\n        0,\n        num_diffusion_steps,\n        (rna.shape[0],),\n        device=device\n    )\n\n    z_noisy, noise = schedule.q_sample(\n        z_rna,\n        t\n    )\n\n    noise_pred = model.atac_to_rna_diffusion(\n        z_noisy,\n        z_atac,\n        t\n    )\n\n    diffusion_loss = F.mse_loss(\n        noise_pred,\n        noise\n    )\n\n    rna_reconstructed = model.multi_rna_decoder(\n        z_rna\n    )\n\n    reconstruction_loss = F.mse_loss(\n        rna_reconstructed,\n        rna\n    )\n\n    loss = (\n        diffusion_loss +\n        0.1 * reconstruction_loss\n    )\n\n    return (\n        loss,\n        diffusion_loss,\n        reconstruction_loss\n    )\n\n\ndef cite_loss(\n    model,\n    rna,\n    protein\n):\n\n    rna = rna.to(\n        device,\n        non_blocking=True\n    )\n\n    protein = protein.to(\n        device,\n        non_blocking=True\n    )\n\n    z_rna = model.encode_cite_rna(rna)\n    z_protein = model.encode_protein(protein)\n\n    t = torch.randint(\n        0,\n        num_diffusion_steps,\n        (protein.shape[0],),\n        device=device\n    )\n\n    z_noisy, noise = schedule.q_sample(\n        z_protein,\n        t\n    )\n\n    noise_pred = model.rna_to_protein_diffusion(\n        z_noisy,\n        z_rna,\n        t\n    )\n\n    diffusion_loss = F.mse_loss(\n        noise_pred,\n        noise\n    )\n\n    protein_reconstructed = model.protein_decoder(\n        z_protein\n    )\n\n    reconstruction_loss = F.mse_loss(\n        protein_reconstructed,\n        protein\n    )\n\n    loss = (\n        diffusion_loss +\n        0.1 * reconstruction_loss\n    )\n\n    return (\n        loss,\n        diffusion_loss,\n        reconstruction_loss\n    )\n\n\nmulti_iterator = itertools.cycle(\n    multi_loader\n)\n\ncite_iterator = itertools.cycle(\n    cite_loader\n)\n\n\ntotal_losses = []\nmultiome_losses = []\ncite_losses = []\n\nbest_loss = float(\"inf\")\nbest_epoch = 0\npatience_counter = 0\n\n\nfor epoch in range(epochs):\n\n    model.train()\n\n    epoch_loss = 0.0\n    epoch_multi = 0.0\n    epoch_cite = 0.0\n\n    steps = max(\n        len(multi_loader),\n        len(cite_loader)\n    )\n\n    for step in range(steps):\n\n        multi_batch = next(\n            multi_iterator\n        )\n\n        cite_batch = next(\n            cite_iterator\n        )\n\n        peak, rna_multi = multi_batch\n        rna_cite, protein = cite_batch\n\n        optimizer.zero_grad(\n            set_to_none=True\n        )\n\n        loss_multi, diff_multi, rec_multi = multiome_loss(\n            model,\n            peak,\n            rna_multi\n        )\n\n        loss_cite, diff_cite, rec_cite = cite_loss(\n            model,\n            rna_cite,\n            protein\n        )\n\n        loss = (\n            loss_multi +\n            loss_cite\n        )\n\n        loss.backward()\n\n        torch.nn.utils.clip_grad_norm_(\n            model.parameters(),\n            1.0\n        )\n\n        optimizer.step()\n\n        epoch_loss += loss.item()\n        epoch_multi += loss_multi.item()\n        epoch_cite += loss_cite.item()\n\n    epoch_loss /= steps\n    epoch_multi /= steps\n    epoch_cite /= steps\n\n    total_losses.append(epoch_loss)\n    multiome_losses.append(epoch_multi)\n    cite_losses.append(epoch_cite)\n\n    print(\n        f\"Epoch {epoch + 1:03d} | \"\n        f\"Total {epoch_loss:.6f} | \"\n        f\"Multiome {epoch_multi:.6f} | \"\n        f\"CITE {epoch_cite:.6f}\"\n    )\n\n    if epoch_loss < best_loss:\n\n        best_loss = epoch_loss\n        best_epoch = epoch + 1\n        patience_counter = 0\n\n        torch.save(\n            model.state_dict(),\n            os.path.join(\n                \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/Diffusion\",\n                \"unified_diffusion_model_best.pt\"\n            )\n        )\n\n    else:\n\n        patience_counter += 1\n\n    if patience_counter >= patience:\n\n        print(\n            f\"Early stopping at epoch {epoch + 1}\"\n        )\n\n        print(\n            f\"Best epoch: {best_epoch}\"\n        )\n\n        print(\n            f\"Best total loss: {best_loss:.6f}\"\n        )\n\n        break\n\n\ntorch.save(\n    model.state_dict(),\n    os.path.join(\n        \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/Diffusion\",\n        \"unified_diffusion_model.pt\"\n    )\n)\n\n\n@torch.no_grad()\ndef sample_rna_latent(\n    model,\n    z_atac\n):\n\n    model.eval()\n\n    batch_size_current = z_atac.shape[0]\n\n    z = torch.randn(\n        batch_size_current,\n        latent_dim,\n        device=device\n    )\n\n    for step in reversed(\n        range(num_diffusion_steps)\n    ):\n\n        t = torch.full(\n            (batch_size_current,),\n            step,\n            device=device,\n            dtype=torch.long\n        )\n\n        beta_t = schedule.beta[step]\n        alpha_t = schedule.alpha[step]\n        alpha_bar_t = schedule.alpha_bar[step]\n\n        noise_pred = model.atac_to_rna_diffusion(\n            z,\n            z_atac,\n            t\n        )\n\n        z = (\n            1.0 / torch.sqrt(alpha_t)\n        ) * (\n            z -\n            (\n                beta_t /\n                torch.sqrt(\n                    1.0 - alpha_bar_t\n                )\n            ) * noise_pred\n        )\n\n        if step > 0:\n\n            noise = torch.randn_like(z)\n\n            z = (\n                z +\n                torch.sqrt(beta_t) * noise\n            )\n\n    return z\n\n\n@torch.no_grad()\ndef sample_protein_latent(\n    model,\n    z_rna\n):\n\n    model.eval()\n\n    batch_size_current = z_rna.shape[0]\n\n    z = torch.randn(\n        batch_size_current,\n        protein_latent_dim,\n        device=device\n    )\n\n    for step in reversed(\n        range(num_diffusion_steps)\n    ):\n\n        t = torch.full(\n            (batch_size_current,),\n            step,\n            device=device,\n            dtype=torch.long\n        )\n\n        beta_t = schedule.beta[step]\n        alpha_t = schedule.alpha[step]\n        alpha_bar_t = schedule.alpha_bar[step]\n\n        noise_pred = model.rna_to_protein_diffusion(\n            z,\n            z_rna,\n            t\n        )\n\n        z = (\n            1.0 / torch.sqrt(alpha_t)\n        ) * (\n            z -\n            (\n                beta_t /\n                torch.sqrt(\n                    1.0 - alpha_bar_t\n                )\n            ) * noise_pred\n        )\n\n        if step > 0:\n\n            noise = torch.randn_like(z)\n\n            z = (\n                z +\n                torch.sqrt(beta_t) * noise\n            )\n\n    return z\n\n\n@torch.no_grad()\ndef predict_rna_from_atac(\n    peak\n):\n\n    if hasattr(peak, \"toarray\"):\n        peak = peak.toarray()\n\n    peak = torch.tensor(\n        np.asarray(\n            peak,\n            dtype=np.float32\n        ),\n        device=device\n    )\n\n    z_atac = model.encode_atac(\n        peak\n    )\n\n    z_rna = sample_rna_latent(\n        model,\n        z_atac\n    )\n\n    rna = model.multi_rna_decoder(\n        z_rna\n    )\n\n    return rna.cpu().numpy()\n\n\n@torch.no_grad()\ndef predict_protein_from_rna(\n    rna\n):\n\n    if hasattr(rna, \"toarray\"):\n        rna = rna.toarray()\n\n    rna = torch.tensor(\n        np.asarray(\n            rna,\n            dtype=np.float32\n        ),\n        device=device\n    )\n\n    z_rna = model.encode_cite_rna(\n        rna\n    )\n\n    z_protein = sample_protein_latent(\n        model,\n        z_rna\n    )\n\n    protein = model.protein_decoder(\n        z_protein\n    )\n\n    return protein.cpu().numpy()\n\n\n@torch.no_grad()\ndef predict_from_atac(\n    peak\n):\n\n    if hasattr(peak, \"toarray\"):\n        peak = peak.toarray()\n\n    peak = torch.tensor(\n        np.asarray(\n            peak,\n            dtype=np.float32\n        ),\n        device=device\n    )\n\n    z_atac = model.encode_atac(\n        peak\n    )\n\n    z_rna = sample_rna_latent(\n        model,\n        z_atac\n    )\n\n    rna = model.multi_rna_decoder(\n        z_rna\n    )\n\n    z_protein = sample_protein_latent(\n        model,\n        z_rna\n    )\n\n    protein = model.protein_decoder(\n        z_protein\n    )\n\n    return (\n        rna.cpu().numpy(),\n        protein.cpu().numpy()\n    )"},{"cell_type":"code","execution_count":11,"id":"9961b74a-3807-4600-8508-8ff0abb1d330","metadata":{},"outputs":[],"source":"import os\nimport matplotlib.pyplot as plt\nimport torch\n\nsave_dir = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/Diffusion\"\nos.makedirs(save_dir, exist_ok=True)\n\ntorch.save(\n    model.state_dict(),\n    os.path.join(save_dir, \"unified_diffusion_model.pt\")\n)\n\nepochs = range(1, 51)\n\ntotal_losses = [\n    0.807554, 0.373368, 0.332057, 0.308909, 0.294314,\n    0.281139, 0.273483, 0.268358, 0.263471, 0.259960,\n    0.257332, 0.254955, 0.253008, 0.251682, 0.249867,\n    0.248614, 0.247936, 0.247016, 0.246345, 0.245966,\n    0.244969, 0.244344, 0.244296, 0.243466, 0.243119,\n    0.242382, 0.242341, 0.241777, 0.241481, 0.240982,\n    0.240541, 0.240589, 0.240313, 0.240028, 0.239808,\n    0.239695, 0.239228, 0.239116, 0.238853, 0.238454,\n    0.238575, 0.238040, 0.238185, 0.237862, 0.237632,\n    0.237694, 0.237413, 0.237325, 0.237317, 0.237112\n]\n\nmultiome_losses = [\n    0.399912, 0.225686, 0.219658, 0.215587, 0.213098,\n    0.208748, 0.206747, 0.205170, 0.204008, 0.202967,\n    0.202140, 0.201194, 0.200676, 0.200428, 0.199500,\n    0.199013, 0.198944, 0.198814, 0.198453, 0.198447,\n    0.197906, 0.197783, 0.197786, 0.197577, 0.197140,\n    0.197196, 0.197235, 0.196765, 0.196805, 0.196609,\n    0.196402, 0.196596, 0.196123, 0.196208, 0.196317,\n    0.196151, 0.195969, 0.195927, 0.195829, 0.195740,\n    0.195714, 0.195576, 0.195401, 0.195467, 0.195311,\n    0.195365, 0.195259, 0.195254, 0.195219, 0.195271\n]\n\ncite_losses = [\n    0.407642, 0.147682, 0.112399, 0.093322, 0.081216,\n    0.072391, 0.066736, 0.063188, 0.059464, 0.056994,\n    0.055192, 0.053761, 0.052332, 0.051254, 0.050368,\n    0.049601, 0.048992, 0.048203, 0.047892, 0.047519,\n    0.047063, 0.046561, 0.046511, 0.045888, 0.045979,\n    0.045187, 0.045106, 0.045012, 0.044676, 0.044373,\n    0.044140, 0.043994, 0.044190, 0.043821, 0.043492,\n    0.043544, 0.043259, 0.043189, 0.043025, 0.042714,\n    0.042861, 0.042465, 0.042784, 0.042396, 0.042322,\n    0.042330, 0.042154, 0.042071, 0.042097, 0.041841\n]\n\nplt.figure(figsize=(10, 8))\n\nplt.plot(\n    epochs,\n    total_losses,\n    marker=\"o\",\n    label=\"Total Loss\"\n)\n\nplt.plot(\n    epochs,\n    multiome_losses,\n    marker=\"o\",\n    label=\"Multiome Loss\"\n)\n\nplt.plot(\n    epochs,\n    cite_losses,\n    marker=\"o\",\n    label=\"CITE Loss\"\n)\n\nplt.xlabel(\"Epoch\")\nplt.ylabel(\"Loss\")\nplt.title(\"Training Loss\")\nplt.xticks(list(epochs))\nplt.legend()\nplt.grid(True, alpha=0.3)\nplt.tight_layout()\n\nplt.savefig(\n    os.path.join(save_dir, \"training_loss.png\"),\n    dpi=300,\n    bbox_inches=\"tight\"\n)\n\nplt.show()"},{"cell_type":"code","execution_count":14,"id":"c10a83e2-eb51-4377-9d2a-0e7cf04b8de8","metadata":{},"outputs":[],"source":"import inspect\n\nprint(inspect.signature(MultimodalDiffusionModel))"},{"cell_type":"code","execution_count":47,"id":"e7387117-8ebc-40be-942c-e52a5cd812b1","metadata":{},"outputs":[],"source":"print(\"RNA\")\nprint(z_rna.mean().item())\nprint(z_rna.std().item())\nprint(z_rna.min().item())\nprint(z_rna.max().item())\n\nprint(\"\\nProtein\")\nprint(z_protein_denoised.mean().item())\nprint(z_protein_denoised.std().item())\nprint(z_protein_denoised.min().item())\nprint(z_protein_denoised.max().item())\n\nrna_var = z_rna.detach().cpu().numpy().var(axis=0)\nprotein_var = z_protein_denoised.detach().cpu().numpy().var(axis=0)\n\nprint(\"RNA latent variance\")\nprint(\"mean:\", rna_var.mean())\nprint(\"median:\", np.median(rna_var))\nprint(\"max:\", rna_var.max())\n\nprint(\"\\nProtein latent variance\")\nprint(\"mean:\", protein_var.mean())\nprint(\"median:\", np.median(protein_var))\nprint(\"max:\", protein_var.max())"},{"cell_type":"code","execution_count":65,"id":"5efa4837-b9bf-4abd-a6db-bab379ef4911","metadata":{},"outputs":[],"source":"import torch\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Patch\n\nwith torch.no_grad():\n    z_rna = model.cite_rna_encoder(rna)\n\n    z_protein = z_rna.clone()\n\n    protein_noise = torch.randn_like(z_protein[..., :64])\n\n    z_protein = z_protein[..., :64] + protein_noise\n\nrna_values = z_rna.cpu().numpy().ravel()\nprotein_noise_values = protein_noise.cpu().numpy().ravel()\n\nplt.figure(figsize=(5, 5))\n\nplt.hist(\n    rna_values,\n    bins=50,\n    density=True,\n    alpha=0.5,\n    color=\"yellow\"\n)\n\nplt.hist(\n    protein_noise_values,\n    bins=50,\n    density=True,\n    alpha=0.5,\n    color=\"red\"\n)\n\nlegend_elements = [\n    Patch(facecolor=\"yellow\", alpha=0.5, label=\"ATAC → RNA\"),\n    Patch(facecolor=\"red\", alpha=0.5, label=\"RNA → Protein\")\n]\n\nplt.xlabel(\"Latent / Noise Value\")\nplt.ylabel(\"Density\")\nplt.title(\"Diffusion Noise latent comparison\")\nplt.legend(handles=legend_elements)\n\nplt.tight_layout()\nplt.show()"},{"cell_type":"code","execution_count":56,"id":"c6113232-a1bc-40c3-8be7-616d06845e8f","metadata":{},"outputs":[],"source":"@torch.no_grad()\ndef sample_protein_latent(model, z_rna):\n    model.eval()\n    batch_size_current = z_rna.shape[0]\n\n    z = torch.randn(\n        batch_size_current,\n        protein_latent_dim,\n        device=device\n    )\n\n    for step in reversed(range(num_diffusion_steps)):\n        t = torch.full(\n            (batch_size_current,),\n            step,\n            device=device,\n            dtype=torch.long\n        )\n\n        beta_t = schedule.beta[step]\n        alpha_t = schedule.alpha[step]\n        alpha_bar_t = schedule.alpha_bar[step]\n\n        noise_pred = model.rna_to_protein_diffusion(\n            z,\n            z_rna,\n            t\n        )\n\n        z = (\n            1.0 / torch.sqrt(alpha_t)\n        ) * (\n            z -\n            (\n                beta_t /\n                torch.sqrt(1.0 - alpha_bar_t)\n            ) * noise_pred\n        )\n\n        if step > 0:\n            noise = torch.randn_like(z)\n            z = z + torch.sqrt(beta_t) * noise\n\n    return z"},{"cell_type":"code","execution_count":66,"id":"45ec55c2-9fd7-4ca3-bb4c-dec5f0019167","metadata":{},"outputs":[],"source":"import torch\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Patch\n\nwith torch.no_grad():\n    if isinstance(peak, torch.Tensor):\n        peak_tensor = peak.to(device)\n    else:\n        if hasattr(peak, \"toarray\"):\n            peak = peak.toarray()\n        peak_tensor = torch.tensor(\n            np.asarray(peak, dtype=np.float32),\n            device=device\n        )\n\n    if isinstance(rna, torch.Tensor):\n        rna_tensor = rna.to(device)\n    else:\n        if hasattr(rna, \"toarray\"):\n            rna = rna.toarray()\n        rna_tensor = torch.tensor(\n            np.asarray(rna, dtype=np.float32),\n            device=device\n        )\n\n    z_atac = model.encode_atac(peak_tensor)\n    z_rna_denoised = sample_rna_latent(model, z_atac)\n\n    z_rna = model.cite_rna_encoder(rna_tensor)\n    z_protein_denoised = sample_protein_latent(model, z_rna)\n\nrna_denoised_values = z_rna_denoised.cpu().numpy().ravel()\nprotein_denoised_values = z_protein_denoised.cpu().numpy().ravel()\n\nplt.figure(figsize=(5, 5))\n\nplt.hist(\n    rna_denoised_values,\n    bins=50,\n    density=True,\n    alpha=0.5,\n    color=\"yellow\"\n)\n\nplt.hist(\n    protein_denoised_values,\n    bins=50,\n    density=True,\n    alpha=0.5,\n    color=\"red\"\n)\n\nlegend_elements = [\n    Patch(facecolor=\"yellow\", alpha=0.5, label=\"ATAC → RNA\"),\n    Patch(facecolor=\"red\", alpha=0.5, label=\"RNA → Protein\")\n]\n\nplt.xlabel(\"Denoised Latent Value\")\nplt.ylabel(\"Density\")\nplt.title(\"Diffusion Denoise latent comparison\")\nplt.legend(handles=legend_elements)\n\nplt.tight_layout()\nplt.show()"},{"cell_type":"code","execution_count":61,"id":"70d38616-9f62-4b49-a27a-7bdc2f55ed21","metadata":{},"outputs":[],"source":"import torch\nimport numpy as np\nimport matplotlib.pyplot as plt\n\n@torch.no_grad()\ndef trace_rna_diffusion(model, z_atac):\n    model.eval()\n\n    batch_size_current = z_atac.shape[0]\n\n    z = torch.randn(\n        batch_size_current,\n        latent_dim,\n        device=device\n    )\n\n    steps = []\n    noise_variances = []\n    denoise_variances = []\n\n    for step in reversed(range(num_diffusion_steps)):\n        t = torch.full(\n            (batch_size_current,),\n            step,\n            device=device,\n            dtype=torch.long\n        )\n\n        beta_t = schedule.beta[step]\n        alpha_t = schedule.alpha[step]\n        alpha_bar_t = schedule.alpha_bar[step]\n\n        noise_pred = model.atac_to_rna_diffusion(\n            z,\n            z_atac,\n            t\n        )\n\n        z_denoised = (\n            1.0 / torch.sqrt(alpha_t)\n        ) * (\n            z -\n            (\n                beta_t /\n                torch.sqrt(1.0 - alpha_bar_t)\n            ) * noise_pred\n        )\n\n        noise_variances.append(\n            z.var(dim=0).mean().item()\n        )\n\n        denoise_variances.append(\n            z_denoised.var(dim=0).mean().item()\n        )\n\n        steps.append(step)\n\n        z = z_denoised\n\n        if step > 0:\n            noise = torch.randn_like(z)\n            z = z + torch.sqrt(beta_t) * noise\n\n    return (\n        np.array(steps),\n        np.array(noise_variances),\n        np.array(denoise_variances)\n    )\n\n\n@torch.no_grad()\ndef trace_protein_diffusion(model, z_rna):\n    model.eval()\n\n    batch_size_current = z_rna.shape[0]\n\n    z = torch.randn(\n        batch_size_current,\n        protein_latent_dim,\n        device=device\n    )\n\n    steps = []\n    noise_variances = []\n    denoise_variances = []\n\n    for step in reversed(range(num_diffusion_steps)):\n        t = torch.full(\n            (batch_size_current,),\n            step,\n            device=device,\n            dtype=torch.long\n        )\n\n        beta_t = schedule.beta[step]\n        alpha_t = schedule.alpha[step]\n        alpha_bar_t = schedule.alpha_bar[step]\n\n        noise_pred = model.rna_to_protein_diffusion(\n            z,\n            z_rna,\n            t\n        )\n\n        z_denoised = (\n            1.0 / torch.sqrt(alpha_t)\n        ) * (\n            z -\n            (\n                beta_t /\n                torch.sqrt(1.0 - alpha_bar_t)\n            ) * noise_pred\n        )\n\n        noise_variances.append(\n            z.var(dim=0).mean().item()\n        )\n\n        denoise_variances.append(\n            z_denoised.var(dim=0).mean().item()\n        )\n\n        steps.append(step)\n\n        z = z_denoised\n\n        if step > 0:\n            noise = torch.randn_like(z)\n            z = z + torch.sqrt(beta_t) * noise\n\n    return (\n        np.array(steps),\n        np.array(noise_variances),\n        np.array(denoise_variances)\n    )\n\n\nwith torch.no_grad():\n    if isinstance(peak, torch.Tensor):\n        peak_tensor = peak.to(device)\n    else:\n        if hasattr(peak, \"toarray\"):\n            peak = peak.toarray()\n        peak_tensor = torch.tensor(\n            np.asarray(peak, dtype=np.float32),\n            device=device\n        )\n\n    if isinstance(rna, torch.Tensor):\n        rna_tensor = rna.to(device)\n    else:\n        if hasattr(rna, \"toarray\"):\n            rna = rna.toarray()\n        rna_tensor = torch.tensor(\n            np.asarray(rna, dtype=np.float32),\n            device=device\n        )\n\n    z_atac = model.encode_atac(peak_tensor)\n    z_rna = model.cite_rna_encoder(rna_tensor)\n\nrna_steps, rna_noise_var, rna_denoise_var = trace_rna_diffusion(\n    model,\n    z_atac\n)\n\nprotein_steps, protein_noise_var, protein_denoise_var = trace_protein_diffusion(\n    model,\n    z_rna\n)"},{"cell_type":"code","execution_count":67,"id":"c7116b4f-527f-431b-a485-ed2317f8b573","metadata":{},"outputs":[],"source":"plt.figure(figsize=(5, 5))\n\nplt.plot(\n    rna_steps,\n    rna_noise_var,\n    label=\"ATAC → RNA\"\n)\n\nplt.plot(\n    protein_steps,\n    protein_noise_var,\n    label=\"RNA → Protein\"\n)\n\nplt.xlabel(\"Diffusion Step\")\nplt.ylabel(\"Noise Variance\")\nplt.title(\"Noise Variance Across Diffusion Steps\")\nplt.legend()\nplt.tight_layout()\nplt.show()"},{"cell_type":"code","execution_count":68,"id":"057afafb-cf30-4fa1-b518-dc10b4a845dd","metadata":{},"outputs":[],"source":"plt.figure(figsize=(5, 5))\n\nplt.plot(\n    rna_steps,\n    rna_denoise_var,\n    label=\"ATAC → RNA\"\n)\n\nplt.plot(\n    protein_steps,\n    protein_denoise_var,\n    label=\"RNA → Protein\"\n)\n\nplt.xlabel(\"Diffusion Step\")\nplt.ylabel(\"Denoised Latent Variance\")\nplt.title(\"Denoised Latent Variance Across Diffusion Steps\")\nplt.gca().invert_xaxis()\nplt.legend()\nplt.tight_layout()\nplt.show()"},{"cell_type":"code","execution_count":78,"id":"566ebea3-4eaa-4cc1-8b1e-854f8c76c880","metadata":{},"outputs":[],"source":"import numpy as np\nimport torch\n\nbatch_size_eval = 128\n\nrna_squared_errors = []\nprotein_squared_errors = []\n\nmodel.eval()\n\nwith torch.no_grad():\n\n    peak_data = train_multi.layers[\"peak_sum\"]\n    rna_gt_data = train_multi.layers[\"rna\"]\n\n    for start in range(0, train_multi.n_obs, batch_size_eval):\n        end = min(start + batch_size_eval, train_multi.n_obs)\n\n        peak_batch = peak_data[start:end]\n        rna_gt_batch = rna_gt_data[start:end]\n\n        if hasattr(peak_batch, \"toarray\"):\n            peak_batch = peak_batch.toarray()\n\n        if hasattr(rna_gt_batch, \"toarray\"):\n            rna_gt_batch = rna_gt_batch.toarray()\n\n        peak_batch = torch.tensor(\n            np.asarray(peak_batch, dtype=np.float32),\n            device=device\n        )\n\n        rna_gt_batch = torch.tensor(\n            np.asarray(rna_gt_batch, dtype=np.float32),\n            device=device\n        )\n\n        z_atac = model.encode_atac(peak_batch)\n        z_rna_pred = sample_rna_latent(model, z_atac)\n        rna_pred = model.multi_rna_decoder(z_rna_pred)\n\n        mse = (\n            (rna_pred - rna_gt_batch) ** 2\n        ).mean(dim=1)\n\n        rna_squared_errors.append(\n            mse.cpu().numpy()\n        )\n\n        del peak_batch, rna_gt_batch\n        del z_atac, z_rna_pred, rna_pred, mse\n\n        torch.cuda.empty_cache()\n\n    cite_rna_data = train_cite_rna_unique.X\n    protein_gt_data = train_cite[\"protein\"].X\n\n    for start in range(0, train_cite_rna_unique.n_obs, batch_size_eval):\n        end = min(start + batch_size_eval, train_cite_rna_unique.n_obs)\n\n        rna_batch = cite_rna_data[start:end]\n        protein_gt_batch = protein_gt_data[start:end]\n\n        if hasattr(rna_batch, \"toarray\"):\n            rna_batch = rna_batch.toarray()\n\n        if hasattr(protein_gt_batch, \"toarray\"):\n            protein_gt_batch = protein_gt_batch.toarray()\n\n        rna_batch = torch.tensor(\n            np.asarray(rna_batch, dtype=np.float32),\n            device=device\n        )\n\n        protein_gt_batch = torch.tensor(\n            np.asarray(protein_gt_batch, dtype=np.float32),\n            device=device\n        )\n\n        z_rna = model.encode_cite_rna(rna_batch)\n        z_protein_pred = sample_protein_latent(model, z_rna)\n        protein_pred = model.protein_decoder(z_protein_pred)\n\n        mse = (\n            (protein_pred - protein_gt_batch) ** 2\n        ).mean(dim=1)\n\n        protein_squared_errors.append(\n            mse.cpu().numpy()\n        )\n\n        del rna_batch, protein_gt_batch\n        del z_rna, z_protein_pred, protein_pred, mse\n\n        torch.cuda.empty_cache()\n\nrna_squared_errors = np.concatenate(rna_squared_errors)\nprotein_squared_errors = np.concatenate(protein_squared_errors)\n\nrna_rmse = np.sqrt(np.mean(rna_squared_errors))\nprotein_rmse = np.sqrt(np.mean(protein_squared_errors))\n\nprint(f\"ATAC → RNA RMSE: {rna_rmse:.6f}\")\nprint(f\"RNA → Protein RMSE: {protein_rmse:.6f}\")"},{"cell_type":"code","execution_count":83,"id":"99b96335-3f3d-4a2a-922d-9c56097cbdab","metadata":{},"outputs":[],"source":"import matplotlib.pyplot as plt\n\nmodels = [\n    \"ATAC → RNA\",\n    \"RNA → Protein\"\n]\n\nrmse_values = [\n    rna_rmse,\n    protein_rmse\n]\n\nfig, ax = plt.subplots(figsize=(5, 5))\n\nbars = ax.bar(\n    models,\n    rmse_values\n)\n\nax.set_ylabel(\"RMSE\")\nax.set_title(\"Model RMSE comparison\")\n\nax.tick_params(\n    axis=\"both\",\n    labelsize=13\n)\nax.set_ylabel(\"RMSE\", fontsize=18)\nfor bar, value in zip(bars, rmse_values):\n    ax.text(\n        bar.get_x() + bar.get_width() / 2,\n        bar.get_height(),\n        f\"{value:.4f}\",\n        ha=\"center\",\n        va=\"bottom\",\n        fontsize=12\n    )\n\nplt.tight_layout()\nplt.show()"},{"cell_type":"code","execution_count":null,"id":"f862f3f1-dee9-47e5-9472-158a7c62e16e","metadata":{},"outputs":[],"source":""},{"cell_type":"markdown","id":"d0268682-e0de-4805-8bf7-851186c22133","metadata":{},"source":"# Diffusion Predict"},{"cell_type":"code","execution_count":12,"id":"ee3c6a78-0b6c-4c66-b306-5d79c8332655","metadata":{},"outputs":[],"source":"@torch.no_grad()\ndef predict_rna_from_atac(model, peak):\n    z_atac = model.atac_encoder(peak)\n\n    z_rna = sample_rna_latent(\n        model,\n        z_atac\n    )\n\n    rna = model.multi_rna_decoder(z_rna)\n\n    return rna\n\n\n@torch.no_grad()\ndef predict_protein_from_rna(model, rna):\n    z_rna = model.cite_rna_encoder(rna)\n\n    z_protein = sample_protein_latent(\n        model,\n        z_rna\n    )\n\n    protein = model.protein_decoder(z_protein)\n\n    return protein"},{"cell_type":"code","execution_count":13,"id":"300684cf-11fd-4a42-8749-09041c49158f","metadata":{},"outputs":[],"source":"import os\nimport numpy as np\nimport torch\nimport anndata as ad\n\ndevice = torch.device(\n    \"cuda:0\" if torch.cuda.is_available() else \"cpu\"\n)\n\nsave_dir = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/Diffusion\"\n\nmodel_path = os.path.join(\n    save_dir,\n    \"unified_diffusion_model.pt\"\n)\n\nmodel.load_state_dict(\n    torch.load(\n        model_path,\n        map_location=device,\n        weights_only=True\n    )\n)\n\nmodel = model.to(device)\nmodel.eval()\n\nprint(\"Model loaded\")\nprint(\n    \"Device:\",\n    next(model.parameters()).device\n)\n\n\ncite_keep = ~train_cite[\"rna\"].var_names.duplicated(\n    keep=\"first\"\n)\n\ntrain_cite_rna_unique = train_cite[\"rna\"][\n    :,\n    cite_keep\n].copy()\n\ntest_rna_keep = ~test_rna.var_names.duplicated(\n    keep=\"first\"\n)\n\ntest_rna = test_rna[\n    :,\n    test_rna_keep\n].copy()\n\n\ntrain_multi_genes = train_multi.var_names.to_numpy()\n\ntrain_cite_genes = train_cite_rna_unique.var_names.to_numpy()\n\n\ntest_atac = test_atac[\n    :,\n    train_multi_genes\n].copy()\n\ntest_rna = test_rna[\n    :,\n    train_cite_genes\n].copy()\n\n\nassert np.array_equal(\n    test_atac.var_names,\n    train_multi.var_names\n)\n\nassert np.array_equal(\n    test_rna.var_names,\n    train_cite_rna_unique.var_names\n)\n\n\nprint(\n    \"ATAC feature order:\",\n    np.array_equal(\n        test_atac.var_names,\n        train_multi.var_names\n    )\n)\n\nprint(\n    \"RNA feature order:\",\n    np.array_equal(\n        test_rna.var_names,\n        train_cite_rna_unique.var_names\n    )\n)\n\nprint(\n    \"Multiome features:\",\n    train_multi.n_vars\n)\n\nprint(\n    \"CITE RNA features:\",\n    train_cite_rna_unique.n_vars\n)\n\nprint(\n    \"Test RNA features:\",\n    test_rna.n_vars\n)\n\n\nbatch_size = 128\n\n\npeak = test_atac.layers[\"peak_sum\"]\n\nif hasattr(peak, \"toarray\"):\n    peak = peak.toarray()\n\npeak = np.asarray(\n    peak,\n    dtype=np.float32\n)\n\n\npred_rna_list = []\n\nwith torch.no_grad():\n\n    for start in range(\n        0,\n        len(test_atac),\n        batch_size\n    ):\n\n        end = min(\n            start + batch_size,\n            len(test_atac)\n        )\n\n        peak_batch = torch.tensor(\n            peak[start:end],\n            dtype=torch.float32,\n            device=device\n        )\n\n        z_atac = model.encode_atac(\n            peak_batch\n        )\n\n        z_rna = sample_rna_latent(\n            model,\n            z_atac\n        )\n\n        pred_rna_batch = model.multi_rna_decoder(\n            z_rna\n        )\n\n        pred_rna_list.append(\n            pred_rna_batch.cpu().numpy()\n        )\n\n\npred_rna = np.concatenate(\n    pred_rna_list,\n    axis=0\n)\n\n\npred_rna_adata = ad.AnnData(\n    X=pred_rna,\n    obs=test_atac.obs.copy(),\n    var=train_multi.var.copy()\n)\n\nprint(\n    \"Predicted RNA:\",\n    pred_rna_adata.shape\n)\n\n\nrna = test_rna.X\n\nif hasattr(rna, \"toarray\"):\n    rna = rna.toarray()\n\nrna = np.asarray(\n    rna,\n    dtype=np.float32\n)\n\n\npred_protein_list = []\n\nwith torch.no_grad():\n\n    for start in range(\n        0,\n        len(test_rna),\n        batch_size\n    ):\n\n        end = min(\n            start + batch_size,\n            len(test_rna)\n        )\n\n        rna_batch = torch.tensor(\n            rna[start:end],\n            dtype=torch.float32,\n            device=device\n        )\n\n        z_rna = model.encode_cite_rna(\n            rna_batch\n        )\n\n        z_protein = sample_protein_latent(\n            model,\n            z_rna\n        )\n\n        pred_protein_batch = model.protein_decoder(\n            z_protein\n        )\n\n        pred_protein_list.append(\n            pred_protein_batch.cpu().numpy()\n        )\n\n\npred_protein = np.concatenate(\n    pred_protein_list,\n    axis=0\n)\n\n\npred_protein_adata = ad.AnnData(\n    X=pred_protein,\n    obs=test_rna.obs.copy(),\n    var=train_cite[\"protein\"].var.copy()\n)\n\n\nprint(\n    \"Predicted Protein:\",\n    pred_protein_adata.shape\n)\n\n\nrna_path = os.path.join(\n    save_dir,\n    \"test_predicted_rna.h5ad\"\n)\n\nprotein_path = os.path.join(\n    save_dir,\n    \"test_predicted_protein.h5ad\"\n)\n\n\npred_rna_adata.write_h5ad(\n    rna_path\n)\n\npred_protein_adata.write_h5ad(\n    protein_path\n)\n\n\nprint(\"Saved:\")\nprint(rna_path)\nprint(protein_path)"},{"cell_type":"code","execution_count":33,"id":"eba83225-4f2f-46a4-91a5-5c6744ac20b1","metadata":{},"outputs":[],"source":"import pandas as pd\n\nevaluation_ids = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/evaluation_ids.csv\"\n)\n\nsample_submission = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/sample_submission.csv\"\n)"},{"cell_type":"code","execution_count":37,"id":"b2fd25fd-c732-4b59-80f1-2c1460e08525","metadata":{},"outputs":[],"source":"import pandas as pd\nimport requests\nimport time\n\nevaluation_ids = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/evaluation_ids.csv\"\n)\n\ngtf_path = \"/home/Data_Drive_8TB/kykim/7. Kaggle/Single-Cell_Perturbations/baseline/gencode.v44.annotation.gtf.gz\"\n\ngtf = pd.read_csv(\n    gtf_path,\n    sep=\"\\t\",\n    comment=\"#\",\n    header=None,\n    names=[\n        \"chr\", \"source\", \"feature\", \"start\", \"end\",\n        \"score\", \"strand\", \"frame\", \"attribute\"\n    ]\n)\n\ngene_gtf = gtf[gtf[\"feature\"] == \"gene\"].copy()\n\ngene_gtf[\"gene_id\"] = (\n    gene_gtf[\"attribute\"]\n    .str.extract(r'gene_id \"([^\"]+)\"')[0]\n    .str.split(\".\")\n    .str[0]\n)\n\ngene_gtf[\"gene_name\"] = (\n    gene_gtf[\"attribute\"]\n    .str.extract(r'gene_name \"([^\"]+)\"')[0]\n)\n\ngene_map = (\n    gene_gtf\n    .dropna(subset=[\"gene_id\", \"gene_name\"])\n    .drop_duplicates(\"gene_id\")\n    .set_index(\"gene_id\")[\"gene_name\"]\n    .to_dict()\n)\n\nmask = evaluation_ids[\"gene_id\"].astype(str).str.match(\n    r\"^ENSG\\d+\",\n    na=False\n)\n\nensg_ids = (\n    evaluation_ids.loc[mask, \"gene_id\"]\n    .astype(str)\n    .str.extract(r\"^(ENSG\\d+)\", expand=False)\n)\n\nmapped_genes = ensg_ids.map(gene_map)\n\nevaluation_ids.loc[mask, \"gene_id\"] = mapped_genes\n\nremaining_mask = (\n    evaluation_ids[\"gene_id\"]\n    .astype(str)\n    .str.match(r\"^ENSG\\d+$\", na=False)\n)\n\nremaining_ids = (\n    evaluation_ids.loc[remaining_mask, \"gene_id\"]\n    .astype(str)\n    .unique()\n)\n\nprint(\"GTF에서 변환되지 않은 ENSG:\", len(remaining_ids))\n\nensembl_map = {}\n\nurl = \"https://rest.ensembl.org/lookup/id\"\n\nfor i in range(0, len(remaining_ids), 1000):\n    batch = remaining_ids[i:i + 1000]\n\n    response = requests.post(\n        url,\n        params={\n            \"species\": \"homo_sapiens\",\n            \"object_type\": \"gene\"\n        },\n        headers={\n            \"Content-Type\": \"application/json\",\n            \"Accept\": \"application/json\"\n        },\n        json={\n            \"ids\": batch.tolist()\n        }\n    )\n\n    response.raise_for_status()\n\n    result = response.json()\n\n    for gene_id, info in result.items():\n        if isinstance(info, dict):\n            gene_name = info.get(\"display_name\")\n\n            if gene_name:\n                ensembl_map[gene_id] = gene_name\n\n    print(\n        f\"{min(i + 1000, len(remaining_ids))} / {len(remaining_ids)}\"\n    )\n\n    time.sleep(0.1)\n\nevaluation_ids.loc[remaining_mask, \"gene_id\"] = (\n    evaluation_ids.loc[remaining_mask, \"gene_id\"]\n    .map(ensembl_map)\n    .fillna(evaluation_ids.loc[remaining_mask, \"gene_id\"])\n)\n\nunmapped = evaluation_ids.loc[\n    evaluation_ids[\"gene_id\"].astype(str).str.match(\n        r\"^ENSG\\d+$\",\n        na=False\n    ),\n    \"gene_id\"\n].unique()\n\nprint(\"최종 변환되지 않은 ENSG:\", len(unmapped))\nprint(unmapped[:20])\n\nprint(evaluation_ids.loc[mask, [\"gene_id\"]].tail())"},{"cell_type":"code","execution_count":43,"id":"ff65b1d3-a1f8-4869-9bcd-42dde9629638","metadata":{},"outputs":[],"source":"cell_ids = evaluation_ids[\"cell_id\"].astype(str)\ngene_ids = evaluation_ids[\"gene_id\"].astype(str)\n\nrna_cells = pd.Series(\n    np.arange(pred_rna_adata.n_obs),\n    index=pred_rna_adata.obs_names.astype(str)\n)\nrna_genes = pd.Series(\n    np.arange(pred_rna_adata.n_vars),\n    index=pred_rna_adata.var_names.astype(str)\n)\n\nprotein_cells = pd.Series(\n    np.arange(pred_protein_adata.n_obs),\n    index=pred_protein_adata.obs_names.astype(str)\n)\nprotein_genes = pd.Series(\n    np.arange(pred_protein_adata.n_vars),\n    index=pred_protein_adata.var_names.astype(str)\n)\n\nrna_r = cell_ids.map(rna_cells)\nrna_c = gene_ids.map(rna_genes)\n\nprotein_r = cell_ids.map(protein_cells)\nprotein_c = gene_ids.map(protein_genes)\n\nrna_valid = rna_r.notna() & rna_c.notna()\nprotein_valid = protein_r.notna() & protein_c.notna()\n\ntarget = np.zeros(len(evaluation_ids), dtype=np.float32)\n\nif rna_valid.any():\n    r = rna_r[rna_valid].to_numpy(dtype=np.int64)\n    c = rna_c[rna_valid].to_numpy(dtype=np.int64)\n    target[rna_valid.to_numpy()] = np.asarray(\n        pred_rna_adata.X[r, c]\n    ).ravel()\n\nif protein_valid.any():\n    protein_only = (~rna_valid) & protein_valid\n    r = protein_r[protein_only].to_numpy(dtype=np.int64)\n    c = protein_c[protein_only].to_numpy(dtype=np.int64)\n    target[protein_only.to_numpy()] = np.asarray(\n        pred_protein_adata.X[r, c]\n    ).ravel()\n\nsample_submission[\"target\"] = target\n\nmatched = rna_valid | protein_valid\ntotal = len(evaluation_ids)\n\nprint(\"전체:\", total)\nprint(\"RNA 매칭:\", rna_valid.sum(), f\"({rna_valid.sum() / total * 100:.2f}%)\")\nprint(\"Protein 매칭:\", protein_valid.sum(), f\"({protein_valid.sum() / total * 100:.2f}%)\")\nprint(\"매칭 성공:\", matched.sum(), f\"({matched.sum() / total * 100:.2f}%)\")\nprint(\"없는 데이터:\", (~matched).sum(), f\"({(~matched).sum() / total * 100:.2f}%)\")"},{"cell_type":"code","execution_count":47,"id":"8fb6ea18-f06b-4e13-88bc-d6b40bd78139","metadata":{},"outputs":[],"source":"eval_genes = set(evaluation_ids[\"gene_id\"].astype(str))\n\nrna_genes_set = set(pred_rna_adata.var_names.astype(str))\nprotein_genes_set = set(pred_protein_adata.var_names.astype(str))\n\nrna_eval_only = sorted(eval_genes - rna_genes_set)\nrna_object_only = sorted(rna_genes_set - eval_genes)\n\nprotein_eval_only = sorted(eval_genes - protein_genes_set)\nprotein_object_only = sorted(protein_genes_set - eval_genes)\n\nprint(\"===== RNA =====\")\nprint(\"evaluation_ids에만 있는 gene:\", len(rna_eval_only))\nprint(rna_eval_only)\n\nprint(\"\\npred_rna_adata에만 있는 gene:\", len(rna_object_only))\nprint(rna_object_only)\n\nprint(\"\\n===== Protein =====\")\nprint(\"evaluation_ids에만 있는 gene:\", len(protein_eval_only))\nprint(protein_eval_only)\n\nprint(\"\\npred_protein_adata에만 있는 gene:\", len(protein_object_only))\nprint(protein_object_only)"},{"cell_type":"code","execution_count":53,"id":"237e98a3-cf1e-4a68-8db1-dda5cfba7c3e","metadata":{},"outputs":[],"source":"sample_submission"},{"cell_type":"code","execution_count":54,"id":"0b342d31-4b02-47a4-b07b-9838e4663863","metadata":{},"outputs":[],"source":"B = pd.read_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/sample_submission.csv\"\n)\n\nB"},{"cell_type":"code","execution_count":44,"id":"733f3813-e45b-40d8-8902-dfcfac867ac1","metadata":{},"outputs":[],"source":"sample_submission"},{"cell_type":"code","execution_count":45,"id":"78761139-1d50-43bd-9f83-66690a361fc4","metadata":{},"outputs":[],"source":"sample_submission.to_csv(\n    \"/home/Data_Drive_8TB/kykim/7. Kaggle/Multimodal_Single-Cell_Integration/submission/sample_submission.csv\",\n    index=False\n)"},{"cell_type":"code","execution_count":52,"id":"1c95d288-a321-4767-9679-939f8fe22ed8","metadata":{},"outputs":[],"source":"sample_submission"},{"cell_type":"code","execution_count":null,"id":"b8e68688-6c2e-4b24-bfed-7dbb88da3ace","metadata":{},"outputs":[],"source":""},{"cell_type":"code","execution_count":null,"id":"1124bcd8-8e57-4c23-8232-8ba90057dcee","metadata":{},"outputs":[],"source":""},{"cell_type":"code","execution_count":null,"id":"1e2ff2f0-ebfa-467b-9680-7f5c69f2a552","metadata":{},"outputs":[],"source":""},{"cell_type":"code","execution_count":null,"id":"8f09bd0c-7d7d-4904-9a7d-a2f73266bda4","metadata":{},"outputs":[],"source":""},{"cell_type":"code","execution_count":null,"id":"e6bfe0a4-7e8d-45a9-a4f4-54a33fb96466","metadata":{},"outputs":[],"source":""},{"cell_type":"code","execution_count":null,"id":"7d2af1ae-13c8-4155-98a3-11f850ddcc1f","metadata":{},"outputs":[],"source":""},{"cell_type":"code","execution_count":null,"id":"2477874e-bbca-4bf8-9436-3c65c62ed0b8","metadata":{},"outputs":[],"source":""},{"cell_type":"code","execution_count":null,"id":"0b440088-4f0a-4533-b3ff-60677f488017","metadata":{},"outputs":[],"source":""},{"cell_type":"code","execution_count":null,"id":"b0c93505-32fe-41d3-b01c-059691db2ac3","metadata":{},"outputs":[],"source":""},{"cell_type":"code","execution_count":null,"id":"fa0ba9e7-302e-4fce-8d00-0d068f40a8b8","metadata":{},"outputs":[],"source":""}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.15"}},"nbformat":4,"nbformat_minor":5}