OSCR

Cross-Species Aging Knowledge Integration into Agentic AI Platform Uncovers Conserved Mechanisms

Code ↔ Paper

21 matches between paragraphs of the paper and lines of its authors' code, computed by the harvester (lexical-v1). Click a colored paragraph or line to see its counterpart.

The 21 matches
  1. [1] § Material and Methods › Data Source Curation and Collection ↔ pipeline/11_ploting/all_figures-raidhani.ipynb, lines 5255–5312 · score 1.00 · Worm Interactome Database, BioGrakn, CROssBAR, DTINet, BindingDB, BioGRID
  2. [2] § Material and Methods › Data Source Curation and Collection ↔ Frontend/components/AboutUs.py, lines 187–246 · score 0.96 · CellAge, GeneAge, MetaboAge, AgeAnno, DrugAge, AgeXtend
  3. [3] § Material and Methods › Knowledge Graph Embedding (KGE) Training ↔ Backend/dgl-ke/python/dglke/utils.py, lines 252–350 · score 0.85 · SimplE, ComplEx, TransE, RotatE, loss function, DGL KE
  4. [4] § Results › EvoAge Enables Accurate Hypothesis Generation and Experimental Validation ↔ Frontend/hypo_agents.py, lines 395–452 · score 0.81 · Amyloid Precursor Protein, secretase enzyme, amyloidogenic processing, APP, compartmental, Postsynapse
  5. [5] § Material and Methods › Knowledge Graph Embedding (KGE) Training ↔ Backend/dgl-ke/python/dglke/eval.py, lines 41–105 · score 0.80 · SimplE, ComplEx, TransE, RotatE, DGL KE, RESCAL
  6. [6] § Results › A Systems-Scale Knowledge Graph Integrating Aging Biology Across Model Species ↔ pipeline/02_data_processing/drugage.ipynb, lines 1–65 · score 0.79 · Saccharomyces cerevisiae, Mus musculus, Danio rerio, Caenorhabditis elegans, Drosophila melanogaster, pipeline
  7. [7] § Results › A Systems-Scale Knowledge Graph Integrating Aging Biology Across Model Species ↔ pipeline/04_orthology_mapping/Run_2_Final_map_orthologs_to_human_desc.R, lines 1–41 · score 0.79 · Saccharomyces cerevisiae, Mus musculus, Danio rerio, Caenorhabditis elegans, Drosophila melanogaster, pipeline
  8. [8] § Material and Methods › Data Preprocessing and Harmonization ↔ pipeline/02_data_processing/tarkg.ipynb, lines 1–97 · score 0.78 · NCBI Gene, UniProt, PubChem, Cellular Component, Reactome, HPO
  9. [9] § Material and Methods › Backend API and Service Architecture ↔ Frontend/evo-utils/src/kani_utils/kani_streamlit_server.py, lines 147–276 · score 0.78 · get_sample_triples, check_relationship, search_biological_entities, server, subgraph, retrieval
  10. [10] § Material and Methods › Backend API and Service Architecture ↔ Frontend/components/MicroServices.py, lines 651–716 · score 0.77 · entity_relationships, check_relationship, sample triples, service, subgraph, Link Prediction
  11. [11] § Material and Methods › Data Preprocessing and Harmonization ↔ pipeline/02_data_processing/PrimeKG.ipynb, lines 1–90 · score 0.77 · NCBI Gene, DrugBank, PubChem, Cellular Component, MONDO, HPO
  12. [12] § Results › Systematic Optimization of Knowledge Graph Embeddings Reveals Architectural Dependencies ↔ Backend/dgl-ke/python/dglke/utils.py, lines 252–350 · score 0.72 · SimplE, ComplEx, TransE, RotatE, loss, RESCAL
  13. [13] § Results › Systematic Optimization of Knowledge Graph Embeddings Reveals Architectural Dependencies ↔ pipeline/11_ploting/all_figures-raidhani.ipynb, lines 1679–1745 · score 0.71 · SimplE, ComplEx, TransE, RotatE, MRR, hit
  14. [14] § Results › A Systems-Scale Knowledge Graph Integrating Aging Biology Across Model Species ↔ Frontend/components/LoginSignup.py, lines 208–282 · score 0.68 · human centric, model organisms, cross species, unified, harmonization, EvoAge
  15. [15] § Results › A Conversational AI Interface Leverages the EvoAge KG for Discovery and Validation ↔ Frontend/components/AboutUs.py, lines 187–246 · score 0.68 · FastAPI, Neo4j, capabilities, powered, Streamlit, KE
  16. [16] § Material and Methods › Orthology-Based Cross-Species Integration ↔ pipeline/04_orthology_mapping/Run_2_Final_map_orthologs_to_human_desc.R, lines 1–41 · score 0.58 · human ortholog mapping, Ensembl, cerevisiae, rerio, musculus, melanogaster
  17. [17] § Results › A Systems-Scale Knowledge Graph Integrating Aging Biology Across Model Species ↔ Frontend/components/LoginSignup.py, lines 208–282 · score 0.58 · biological networks, human centric, connectivity, discovery, EvoAge, Biology
  18. [18] § Results › A Conversational AI Interface Leverages the EvoAge KG for Discovery and Validation ↔ Frontend/agents.py, lines 144–204 · score 0.58 · Search Biological Entity, triple score, invokes, Agent, validates, ranked
  19. [19] § Material and Methods › Graph Database Implementation ↔ pipeline/09_evoage_vs_other/evoage_vs_biochat_escargot/Biochatter/generate_schema_info.py, lines 46–137 · score 0.55 · graph database, Neo4j, APOC, Cypher, Knowledge Graph, edges
  20. [20] § Results › A Systems-Scale Knowledge Graph Integrating Aging Biology Across Model Species ↔ pipeline/02_data_processing/ageannomo.ipynb, lines 171–232 · score 0.54 · Danio rerio, Drosophila melanogaster, genomic, schema, mouse, nodes
  21. [21] § Results › EvoAge Enables Accurate Hypothesis Generation and Experimental Validation ↔ Frontend/hypo_agents.py, lines 395–452 · score 0.50 · amyloidogenic processing, endocytic, secretase, postsynaptic, BACE1, phenotype

Paper

Loaded from Europe PMC by your browser, not stored by OSCR: doi.org · Europe PMC

The paper is loaded when this pane is shown.

The authors' code

Jupyter notebook · 5,312 lines · 199 KB · MIT · 2 matches

  1. # %%
  2. import pandas as pd
  3. # %% [markdown]
  4. # ## Main Figure 1
  5. # %% [markdown]
  6. # ## 1b
  7. # %%
  8. """
  9. EvoAge KG — Figure 1b: single-ring sunburst/donut of data sources,
  10. wedge size = log10(Total edges), colored by KG_Type category.
  11. Reads figure1b.csv (columns: Source, KG_Type, Total, log) and reproduces
  12. evoage_sunburst_1b.svg. Output SVG text stays fully editable (svg.fonttype='none').
  13. """
  14. import pandas as pd
  15. import matplotlib.pyplot as plt
  16. import matplotlib as mpl
  17. mpl.rcParams['font.family'] = 'DejaVu Sans'
  18. # ----------------------------------------------------------------------
  19. # Editable-text SVG output (Rai's standing preference)
  20. # ----------------------------------------------------------------------
  21. mpl.rcParams['svg.fonttype'] = 'none'
  22. # ----------------------------------------------------------------------
  23. # Config
  24. # ----------------------------------------------------------------------
  25. CSV_PATH = "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/figure1b.csv" # update path as needed, e.g. server path
  26. OUT_PATH = "/storage/Arushi/090526_EvoAge/kg_formation/all_figures/FIG1/evoage_sunburst_1b_regen.svg"
  27. CATEGORY_ORDER = ["Aging", "Biomedical", "Species connection"]
  28. CATEGORY_COLORS = {
  29. "Aging": "#e8b4b8", # salmon/pink
  30. "Biomedical": "#7fb3b8", # teal
  31. "Species connection": "#f1c88b", # gold
  32. }
  33. DONUT_WIDTH = 0.55 # ring thickness as fraction of radius (donut hole size)
  34. WEDGE_EDGE_COLOR = "white"
  35. WEDGE_EDGE_WIDTH = 1.0
  36. START_ANGLE = 90 # start at 12 o'clock
  37. CLOCKWISE = True # counterclockwise=False in pie()
  38. SHOW_RAW_NUMBERS = True # print each source's raw Total on its wedge
  39. RAW_LABEL_FONTSIZE = 6
  40. RAW_LABEL_MIN_PCT = 0.0 # raise this (e.g. 0.5) to hide labels on tiny wedges
  41. def format_raw(x: float) -> str:
  42. """Human-readable raw count, e.g. 1055910438 -> '1.06B'."""
  43. if x >= 1e9:
  44. return f"{x/1e9:.2f}B"
  45. if x >= 1e6:
  46. return f"{x/1e6:.1f}M"
  47. if x >= 1e3:
  48. return f"{x/1e3:.1f}K"
  49. return f"{int(x)}"
  50. # ----------------------------------------------------------------------
  51. # Load + prep data
  52. # ----------------------------------------------------------------------
  53. df = pd.read_csv(CSV_PATH)
  54. # clean thousands-separator formatted totals (e.g. "4,71,00,000")
  55. df["Total"] = df["Total"].astype(str).str.replace(",", "").astype(float)
  56. # category order = alphabetical (Aging, Biomedical, Species connection);
  57. # within each category, rows are already sorted descending by 'log' —
  58. # preserve that order rather than re-sorting, so ties/order match the source file
  59. df["KG_Type"] = pd.Categorical(df["KG_Type"], categories=CATEGORY_ORDER, ordered=True)
  60. df = df.sort_values(["KG_Type"], kind="stable").reset_index(drop=True)
  61. sizes = df["log"].values # <-- wedge size uses log values
  62. colors = df["KG_Type"].map(CATEGORY_COLORS).values
  63. grand_total = df["Total"].sum()
  64. # ----------------------------------------------------------------------
  65. # Plot
  66. # ----------------------------------------------------------------------
  67. fig, ax = plt.subplots(figsize=(9.57, 11.27)) # matches ~689x811pt canvas
  68. # autopct only receives a percentage, so use a row counter (pie() calls it
  69. # once per wedge, in the same order as `sizes`/`df`) to map back to the
  70. # actual raw Total for that source.
  71. raw_values = df["Total"].tolist()
  72. _row = {"i": 0}
  73. def raw_number_label(pct):
  74. i = _row["i"]
  75. _row["i"] += 1
  76. if not SHOW_RAW_NUMBERS or pct < RAW_LABEL_MIN_PCT:
  77. return ""
  78. return format_raw(raw_values[i])
  79. wedges, _, raw_labels = ax.pie(
  80. sizes,
  81. colors=colors,
  82. startangle=START_ANGLE,
  83. counterclock=not CLOCKWISE,
  84. radius=1.0,
  85. wedgeprops=dict(
  86. width=DONUT_WIDTH,
  87. edgecolor=WEDGE_EDGE_COLOR,
  88. linewidth=WEDGE_EDGE_WIDTH,
  89. ),
  90. autopct=raw_number_label,
  91. pctdistance=1.0 - DONUT_WIDTH / 2, # center of the donut ring
  92. )
  93. for lbl in raw_labels:
  94. lbl.set_fontsize(RAW_LABEL_FONTSIZE)
  95. lbl.set_ha("center")
  96. lbl.set_va("center")
  97. ax.set_aspect("equal")
  98. # ----------------------------------------------------------------------
  99. # Center annotation
  100. # ----------------------------------------------------------------------
  101. ax.text(
  102. 0, 0.12, "1",
  103. ha="center", va="center", fontsize=22, fontweight="bold",
  104. )
  105. ax.text(
  106. 0, -0.02, f"~{grand_total/1e9:.2f} billion",
  107. ha="center", va="center", fontsize=13,
  108. )
  109. ax.text(
  110. 0, -0.14, "head ←|||◀ tail",
  111. ha="center", va="center", fontsize=11,
  112. )
  113. # ----------------------------------------------------------------------
  114. # Title + legend
  115. # ----------------------------------------------------------------------
  116. ax.set_title("EvoAge", fontsize=18, pad=20)
  117. legend_handles = [
  118. plt.Rectangle((0, 0), 1, 1, color=CATEGORY_COLORS[cat])
  119. for cat in CATEGORY_ORDER
  120. ]
  121. ax.legend(
  122. legend_handles,
  123. CATEGORY_ORDER,
  124. loc="upper center",
  125. bbox_to_anchor=(0.5, 0.02),
  126. ncol=3,
  127. frameon=False,
  128. fontsize=10,
  129. )
  130. fig.tight_layout()
  131. fig.savefig(OUT_PATH, format="svg")
  132. print(f"Saved: {OUT_PATH}")
  133. # %% [markdown]
  134. # ## 1c
  135. # %%
  136. """
  137. Vertical grouped bar charts comparing gene counts at the "121" vs
  138. "121_12M" ortholog-mapping stage, per species, for each KG resource.
  139. Plot 1: Aging KG vs Biomedical KG
  140. Plot 2: Aging KG vs EvoAge KG vs Biomedical KG
  141. Input : combined_gene_counts.csv
  142. columns -> Species, Aging_121, Aging_121_12M,
  143. Biomedical_121, Biomedical_121_12M,
  144. EvoAge_121, EvoAge_121_12M
  145. Output: SVG (+PNG) files, vector, publication-ready.
  146. """
  147. import numpy as np
  148. import pandas as pd
  149. import matplotlib.pyplot as plt
  150. from matplotlib.patches import Patch
  151. # ----------------------------------------------------------------------
  152. # 0. Paths -- EDIT THESE to match your local setup
  153. # ----------------------------------------------------------------------
  154. INPUT_CSV = "combined_gene_counts_no_human.csv"
  155. OUT_DIR = "." # where the SVG/PNG figures will be written
  156. # ----------------------------------------------------------------------
  157. # 1. Style
  158. # ----------------------------------------------------------------------
  159. plt.rcParams.update({
  160. "font.family": "sans-serif",
  161. "font.size": 11,
  162. "axes.spines.top": False,
  163. "axes.spines.right": False,
  164. "svg.fonttype": "none", # keep text editable in Illustrator/Inkscape
  165. })
  166. COLOR_121 = "#3F6FA6" # blue -> "121" (before)
  167. COLOR_121_12M = "#C1440E" # orange -> "121_12M" (after)
  168. # established dataset accent colors (used for the small dataset tick labels)
  169. DATASET_COLORS = {
  170. "Aging": "#F08080",
  171. "Biomedical": "#E8A33D",
  172. "EvoAge": "#4C9AAE",
  173. }
  174. def plot_gene_count_bars(df, datasets, title, outfile,
  175. bar_width=0.38, ds_gap=0.25, group_gap=1.1,
  176. figsize_per_col=1.15):
  177. """
  178. df : the combined_gene_counts dataframe
  179. datasets : list of prefixes, e.g. ["Aging", "Biomedical"]
  180. or ["Aging", "EvoAge", "Biomedical"]
  181. Each dataset contributes one pair of bars (121 / 121_12M) per species.
  182. """
  183. species_list = df["Species"].tolist()
  184. n_ds = len(datasets)
  185. fig_w = max(7, figsize_per_col * n_ds * len(species_list) + 2)
  186. fig, ax = plt.subplots(figsize=(fig_w, 6.2))
  187. x_cursor = 0.0
  188. group_centers, group_labels = [], []
  189. tick_pos, tick_labels, tick_colors = [], [], []
  190. for sp in species_list:
  191. row = df[df["Species"] == sp].iloc[0]
  192. group_start = x_cursor
  193. for ds in datasets:
  194. y0 = row[f"{ds}_121"]
  195. y1 = row[f"{ds}_121_12M"]
  196. x_center = x_cursor
  197. if pd.notna(y0) and pd.notna(y1):
  198. ax.bar(x_center - bar_width / 2, y0, width=bar_width,
  199. color=COLOR_121, edgecolor="white", linewidth=0.6, zorder=3)
  200. ax.bar(x_center + bar_width / 2, y1, width=bar_width,
  201. color=COLOR_121_12M, edgecolor="white", linewidth=0.6, zorder=3)
  202. else:
  203. ax.text(x_center, ax.get_ylim()[1] if ax.get_ylim()[1] > 0 else 1000,
  204. "n/a", ha="center", va="bottom", fontsize=8,
  205. color="grey", style="italic")
  206. tick_pos.append(x_center)
  207. tick_labels.append(ds)
  208. tick_colors.append(DATASET_COLORS.get(ds, "black"))
  209. x_cursor += (bar_width * 2 + ds_gap)
  210. group_centers.append((group_start + x_cursor - (bar_width * 2 + ds_gap)) / 2)
  211. group_labels.append(sp)
  212. # light vertical separator between species groups
  213. if sp != species_list[-1]:
  214. sep_x = x_cursor - ds_gap / 2 + group_gap / 2
  215. ax.axvline(sep_x, color="#DDDDDD", lw=0.8, zorder=0)
  216. x_cursor += group_gap
  217. # dataset-level ticks (small, colored)
  218. ax.set_xticks(tick_pos)
  219. ax.set_xticklabels(tick_labels, rotation=90 if n_ds > 2 else 0, fontsize=8)
  220. for tick, c in zip(ax.get_xticklabels(), tick_colors):
  221. tick.set_color(c)
  222. # species-level labels below the dataset ticks
  223. for xc, sp in zip(group_centers, group_labels):
  224. ax.text(xc, -0.13, sp, transform=ax.get_xaxis_transform(),
  225. ha="center", va="top", fontsize=11.5, fontweight="bold")
  226. ax.set_xlim(-bar_width * 1.5, x_cursor - group_gap + bar_width * 1.5)
  227. ax.set_ylabel("Gene count", fontsize=12)
  228. ax.set_title(title, fontsize=13.5, fontweight="bold", pad=14)
  229. ax.margins(y=0.08)
  230. ax.grid(axis="y", color="#EAEAEA", lw=0.8, zorder=0)
  231. legend_elems = [
  232. Patch(facecolor=COLOR_121, label="121"),
  233. Patch(facecolor=COLOR_121_12M, label="121_12M"),
  234. ]
  235. ax.legend(handles=legend_elems, frameon=False, loc="upper left",
  236. bbox_to_anchor=(1.01, 1.0), title="Ortholog stage")
  237. fig.tight_layout()
  238. fig.savefig(f"{OUT_DIR}/{outfile}.svg", bbox_inches="tight")
  239. fig.savefig(f"{OUT_DIR}/{outfile}.png", dpi=300, bbox_inches="tight")
  240. plt.show()
  241. plt.close(fig)
  242. print(f"Saved {outfile}.svg / .png")
  243. if __name__ == "__main__":
  244. df = pd.read_csv(INPUT_CSV)
  245. # Plot 1: Aging vs Biomedical
  246. plot_gene_count_bars(
  247. df,
  248. datasets=["Aging", "Biomedical"],
  249. title="Gene Count Comparison: 1-to-1 vs 1-to-Many orthologs (Aging vs Biomedical KG)",
  250. outfile="genecount_bar_aging_biomedical",
  251. )
  252. # Plot 2: Aging vs EvoAge vs Biomedical
  253. plot_gene_count_bars(
  254. df,
  255. datasets=["Aging", "EvoAge", "Biomedical"],
  256. title="Gene Count Comparison: 1-to1 vs 1-to-1+1-to-M (Aging vs EvoAge vs Biomedical KG)",
  257. outfile="genecount_bar_aging_evoage_biomedical",
  258. )
  259. # %% [markdown]
  260. # ## 1d
  261. # %%
  262. import plotly.io as pio
  263. pio.renderers.default = "iframe"
  264. # %%
  265. """
  266. KG Node-Type Hierarchy Visualization (Sunburst ONLY)
  267. =====================================================
  268. Builds a multi-layer sunburst chart (Dataset -> Species -> Node type)
  269. for your Aging vs Biomedical knowledge graphs, for the node scheme:
  270. - Node (1:1 + 1:N)
  271. Chart details:
  272. - Two top-level branches: Aging (coral/red family) vs Biomedical
  273. (teal family -- same teal used previously for EvoAge) -- children
  274. inherit shades of their branch color.
  275. - Wedge SIZE is driven by log10(count+1), so small node types (e.g.
  276. Chemical, Protein) stay readable next to huge ones (PMID, Mutation)
  277. instead of being squeezed to an invisible sliver.
  278. - Labels and hover text still show the REAL, un-logged count (e.g.
  279. "Chemical 434,768"), not the log value -- log scale only affects size.
  280. - Combines all species into one figure per scheme (species is just one
  281. ring/layer of the hierarchy, so you still see per-species breakdown
  282. without needing six separate panels).
  283. Requirements:
  284. pip install pandas plotly kaleido
  285. SVG export (kaleido) note: recent kaleido versions (v1+) need a local
  286. Chrome install. If `write_image` fails, run this once in your terminal:
  287. plotly_get_chrome
  288. or pin an older, self-contained kaleido that needs no Chrome:
  289. pip install "kaleido==0.2.1"
  290. HOW TO USE
  291. ----------
  292. 1. Edit the FILES dict below so the paths point to your own CSVs.
  293. 2. Run: python kg_hierarchy_sunburst_aging_biomedical.py
  294. 3. Each figure is saved directly as a .svg file next to this script.
  295. (No HTML files are written -- SVG only, per your request.)
  296. """
  297. import numpy as np
  298. import pandas as pd
  299. import plotly.express as px
  300. # ---------------------------------------------------------------------------
  301. # 1. EDIT THESE PATHS to point at your own CSV files
  302. # ---------------------------------------------------------------------------
  303. FILES = {
  304. "1to1_1toN": {
  305. "Aging": "/storage/Arushi/090526_EvoAge/kg_formation/species_wise_kg_building/aging_121_12M_Species_NodeType_unique_counts.csv",
  306. "Biomedical": "/storage/Arushi/090526_EvoAge/kg_formation/species_wise_kg_building/Species_Biomedical_121_12M_NodeType_unique_counts.csv",
  307. },
  308. }
  309. SCHEME_TITLE = {
  310. "1to1_1toN": "Node (1:1 + 1:N)",
  311. }
  312. # Node-type row names (as they appear in the CSV) -> short display labels
  313. NODE_MAP = {
  314. "CellularComponent": "Cellular",
  315. "Disease": "Disease",
  316. "Phenotype": "Phenotype",
  317. "BiologicalProcess": "Biological",
  318. "PlantSpecies": "Plant",
  319. "MolecularFunction": "Molecular",
  320. "Pathway": "Pathway",
  321. "Mutation": "Mutation",
  322. "ChemicalEntity": "Chemical",
  323. "Protein": "Protein",
  324. "Mirna": "miRNA",
  325. "AnatomicalEntity": "Anatomy",
  326. "Gene": "Gene",
  327. "Tissue": "Tissue",
  328. "PMID": "PMID",
  329. }
  330. # Distinct color families for the two datasets (kept for reference/legend)
  331. # NOTE: Biomedical keeps the SAME teal family previously used for EvoAge.
  332. DATASET_COLORS = {
  333. "Aging": "#E85C6B", # coral/red family
  334. "Biomedical": "#1F7A8C", # teal family (same as former EvoAge)
  335. }
  336. # Distinct designated color per node type (leaf-level segments)
  337. NODE_COLORS = {
  338. "Cellular": "#F5E1A4",
  339. "Disease": "#B79FCB",
  340. "Phenotype": "#E0559B",
  341. "Biological": "#7FB8D6",
  342. "Plant": "#5B2A6E",
  343. "Molecular": "#D98E3F",
  344. "Pathway": "#7A4B32",
  345. "Mutation": "#9B9B9B",
  346. "Chemical": "#C9B458",
  347. "Protein": "#F0BFC8",
  348. "miRNA": "#8E4E8F",
  349. "Anatomy": "#E08D6D",
  350. "Gene": "#4472A8",
  351. "Tissue": "#E8A96B",
  352. "PMID": "#B7C88A",
  353. }
  354. # ---------------------------------------------------------------------------
  355. # 2. Reshape each CSV (wide: NodeType x Species) into long hierarchical rows
  356. # ---------------------------------------------------------------------------
  357. def load_long(dataset_name, path):
  358. df = pd.read_csv(path)
  359. df = df[df["NodeType"].isin(NODE_MAP.keys())].copy()
  360. df["NodeType"] = df["NodeType"].map(NODE_MAP)
  361. species_cols = [c for c in df.columns if c != "NodeType"]
  362. long_df = df.melt(id_vars="NodeType", value_vars=species_cols,
  363. var_name="Species", value_name="Count")
  364. long_df["Dataset"] = dataset_name
  365. long_df = long_df[long_df["Count"] > 0] # drop empty branches (cleaner chart)
  366. long_df["LogCount"] = np.log10(long_df["Count"] + 1) # drives wedge SIZE only
  367. return long_df[["Dataset", "Species", "NodeType", "Count", "LogCount"]]
  368. def build_hierarchy(scheme_key):
  369. files = FILES[scheme_key]
  370. frames = [load_long(ds_name, path) for ds_name, path in files.items()]
  371. return pd.concat(frames, ignore_index=True)
  372. def apply_node_colors(fig):
  373. """Give each NodeType leaf its designated color; Dataset/Species parent
  374. rings get a neutral gray so they don't inherit a random default color."""
  375. colors = [NODE_COLORS.get(lab, "#EDEDED") for lab in fig.data[0].labels]
  376. fig.data[0].marker.colors = colors
  377. return fig
  378. # ---------------------------------------------------------------------------
  379. # 3. Chart builder (Sunburst only)
  380. # ---------------------------------------------------------------------------
  381. def make_sunburst(df, title):
  382. fig = px.sunburst(
  383. df,
  384. path=["Dataset", "Species", "NodeType"],
  385. values="LogCount",
  386. branchvalues="total",
  387. custom_data=["Count"],
  388. title=title,
  389. )
  390. fig = apply_node_colors(fig)
  391. fig.update_traces(
  392. texttemplate="%{label}<br>%{customdata[0]:,}",
  393. hovertemplate="%{label}<br>Count: %{customdata[0]:,}<extra></extra>",
  394. insidetextorientation="radial",
  395. marker=dict(line=dict(color="white", width=1.5)),
  396. )
  397. fig.update_layout(
  398. font=dict(family="Arial, sans-serif", size=15),
  399. title=dict(x=0.5, font=dict(size=22, family="Arial, sans-serif")),
  400. margin=dict(t=80, l=10, r=10, b=10),
  401. paper_bgcolor="white",
  402. )
  403. return fig
  404. # ---------------------------------------------------------------------------
  405. # 4. Run for both schemes, sunburst only -- SVG output only
  406. # ---------------------------------------------------------------------------
  407. if __name__ == "__main__":
  408. for scheme_key, scheme_label in SCHEME_TITLE.items():
  409. df = build_hierarchy(scheme_key)
  410. sun_fig = make_sunburst(df, f"Knowledge-graph node types")
  411. try:
  412. sun_fig.write_image(f"sunburst_aging_biomedical_{scheme_key}.svg", width=1400, height=1400, scale=2)
  413. print(f"Saved: sunburst_aging_biomedical_{scheme_key}.svg")
  414. except Exception as e:
  415. print("SVG saved")
  416. sun_fig.show()
  417. # %%
  418. import numpy as np
  419. import pandas as pd
  420. import matplotlib.pyplot as plt
  421. import matplotlib.patches as mpatches
  422. import matplotlib.patheffects as pe
  423. # Keep SVG text as real, editable <text> elements instead of converting every
  424. # label/title/legend entry into vector path outlines (matplotlib's default,
  425. # "svg.fonttype" = "path", bakes text into shapes that can't be edited or
  426. # re-fonted in Illustrator/Inkscape). "none" keeps it as selectable/editable text.
  427. plt.rcParams["svg.fonttype"] = "none"
  428. # -----------------------------------------------------------------------------
  429. # 1. LOAD & PREP
  430. # -----------------------------------------------------------------------------
  431. # Aging: long-format file, already has an edge_type + dataset column
  432. # (same source file as before, just keep the "Aging" rows this time).
  433. aging_path = "/storage/Arushi/090526_EvoAge/kg_formation/species_wise_kg_building/Edgetype.csv"
  434. # Biomedical: wide-format file (NodeType/Relation x Species columns), needs
  435. # reshaping to long + an edge_type derived from the Relation name.
  436. biomedical_path = "/storage/Arushi/090526_EvoAge/kg_formation/species_wise_kg_building/Species_Biomedical_121_12M_RelationType_unique_counts.csv"
  437. EDGE_TYPES_IN_NAME = [
  438. "NegativelyAssociatedWith",
  439. "PositivelyAssociatedWith",
  440. "NotAssociatedWith",
  441. "NoEffect",
  442. "Promotes",
  443. "Inhibits",
  444. ]
  445. def infer_edge_type(relation):
  446. for et in EDGE_TYPES_IN_NAME:
  447. if et in relation:
  448. return et
  449. return "relation"
  450. # --- Aging (long format already) ---
  451. df_aging = pd.read_csv(aging_path)
  452. df_aging = df_aging[df_aging["dataset"] == "Aging"].copy()
  453. df_aging = df_aging[df_aging["edge_type"] != "AssociatedWith"]
  454. df_aging = df_aging[["Relation", "Species", "count", "edge_type"]]
  455. df_aging["dataset"] = "Aging"
  456. # --- Biomedical (wide format -> long) ---
  457. bio_wide = pd.read_csv(biomedical_path)
  458. bio_wide = bio_wide[~bio_wide["Relation"].isin(["Species_AssociatedWith_Nodes", "Total Triples"])].copy()
  459. species_cols = [c for c in bio_wide.columns if c != "Relation"]
  460. df_biomedical = bio_wide.melt(id_vars="Relation", value_vars=species_cols,
  461. var_name="Species", value_name="count")
  462. df_biomedical = df_biomedical[df_biomedical["count"] > 0].copy()
  463. df_biomedical["edge_type"] = df_biomedical["Relation"].apply(infer_edge_type)
  464. df_biomedical["dataset"] = "Biomedical"
  465. # --- combine ---
  466. df = pd.concat([df_aging, df_biomedical], ignore_index=True)
  467. # True log10(count) needs count > 0 (log10(0) is undefined) -- drop any
  468. # leftover zero/negative-count rows before taking logs anywhere below.
  469. df = df[df["count"] > 0].copy()
  470. # IMPORTANT: several raw Relation rows collapse into the same wedge (e.g. many
  471. # distinct relation names all map to edge_type "relation"). Sum their raw
  472. # counts into ONE number per (dataset, Species, edge_type) wedge FIRST, then
  473. # take a single log10 of that sum. Taking log10 of each row separately and
  474. # summing those logs would instead compute log10(count_1 * count_2 * ...) --
  475. # the log of a PRODUCT, not the log of the total -- which makes wedges with
  476. # many small rows balloon past wedges with one genuinely large count. This
  477. # way, wedge size is proportional to log10(true total count) for that wedge.
  478. df = (df.groupby(["dataset", "Species", "edge_type"], sort=False)["count"]
  479. .sum().reset_index())
  480. def format_count(n):
  481. if n >= 1_000_000_000:
  482. return f"{n/1_000_000_000:.2f}B"
  483. elif n >= 1_000_000:
  484. return f"{n/1_000_000:.2f}M"
  485. elif n >= 1_000:
  486. return f"{n/1_000:.1f}k"
  487. return str(int(n))
  488. # -----------------------------------------------------------------------------
  489. # 2. COLOR MAPS
  490. # -----------------------------------------------------------------------------
  491. # NOTE: Biomedical keeps the SAME teal previously used for EvoAge.
  492. dataset_colors = {"Aging": "#f4c2c2", "Biomedical": "#7fb3b8"}
  493. species_colors = {sp: "#e8e8e8" for sp in
  494. ["Human", "Yeast", "Celegans", "Drosophila", "Zebrafish", "Mouse"]}
  495. edge_colors = {
  496. "NoEffect": "#a8546b",
  497. "relation": "#e8ddc7",
  498. "Promotes": "#8ecae6",
  499. "PositivelyAssociatedWith": "#a3c98f",
  500. "NotAssociatedWith": "#b0b0b0",
  501. "NegativelyAssociatedWith": "#b39ddb",
  502. "Inhibits": "#e8a48c",
  503. # AssociatedWith intentionally excluded
  504. }
  505. edge_legend_labels = {
  506. "NoEffect": "no-effect",
  507. "relation": "relation",
  508. "Promotes": "promotes",
  509. "PositivelyAssociatedWith": "positively-associated",
  510. "NotAssociatedWith": "not-associated",
  511. "NegativelyAssociatedWith": "negatively-associated",
  512. "Inhibits": "inhibits",
  513. }
  514. # -----------------------------------------------------------------------------
  515. # 3. ANGLE ASSIGNMENT (recursive, log10 of the TRUE raw total at every level)
  516. # At each level (dataset -> Species -> edge_type), a node's weight is
  517. # log10(sum of its own raw counts) -- recomputed fresh from "count" every
  518. # time, not inherited by summing children's already-logged weights. That
  519. # means a dataset with a few huge counts (e.g. Biomedical, billions) will
  520. # correctly get a bigger wedge than one with many small/medium counts
  521. # (e.g. Aging), because we're logging the real total once, not summing
  522. # several small logs together (which silently favors "many small values").
  523. # -----------------------------------------------------------------------------
  524. def assign_angles(data, group_cols, start_angle=0, end_angle=360):
  525. if not group_cols:
  526. return data
  527. col = group_cols[0]
  528. # Weight of each child at THIS level = log10(sum of raw counts under it).
  529. # Recomputing from raw "count" at every level (rather than summing
  530. # already-logged child weights) is what keeps parent wedges honestly
  531. # proportional to what's actually inside them -- a dataset with a few
  532. # huge counts will correctly out-weigh one with many small counts.
  533. sums = data.groupby(col, sort=False)["count"].sum()
  534. weights = np.log10(sums)
  535. total_weight = weights.sum()
  536. order = weights.sort_values(ascending=False).index
  537. angle = start_angle
  538. frames = []
  539. for key in order:
  540. w = weights[key]
  541. span = (w / total_weight) * (end_angle - start_angle) if total_weight > 0 else 0
  542. sub = data[data[col] == key].copy()
  543. sub = assign_angles(sub, group_cols[1:], angle, angle + span)
  544. sub[f"{col}_start"] = angle
  545. sub[f"{col}_end"] = angle + span
  546. frames.append(sub)
  547. angle += span
  548. return pd.concat(frames)
  549. # -----------------------------------------------------------------------------
  550. # 4. DRAWING HELPERS
  551. # -----------------------------------------------------------------------------
  552. r_hole = 0.28
  553. r_dataset = (r_hole, 0.55)
  554. r_species = (0.58, 0.78)
  555. r_edge = (0.81, 1.05)
  556. r_label = 1.09
  557. def draw_wedge(ax, t1, t2, r_inner, r_outer, color, edgecolor="white", lw=0.8):
  558. if t2 - t1 <= 0:
  559. return
  560. w = mpatches.Wedge((0, 0), r_outer, t1, t2, width=r_outer - r_inner,
  561. facecolor=color, edgecolor=edgecolor, linewidth=lw)
  562. ax.add_patch(w)
  563. def draw_sunburst(ax, df, title):
  564. # ── Ring 1: dataset ──────────────────────────────────────────────────────
  565. for dataset, grp in df.groupby("dataset", sort=False):
  566. t1 = grp["dataset_start"].iloc[0]
  567. t2 = grp["dataset_end"].iloc[0]
  568. draw_wedge(ax, t1, t2, r_dataset[0], r_dataset[1],
  569. dataset_colors.get(dataset, "#cccccc"))
  570. mid = np.deg2rad((t1 + t2) / 2)
  571. r_t = (r_dataset[0] + r_dataset[1]) / 2
  572. ax.text(r_t * np.cos(mid), r_t * np.sin(mid), dataset,
  573. ha="center", va="center", fontsize=13, fontweight="bold")
  574. # ── Ring 2: species ───────────────────────────────────────────────────────
  575. for (dataset, sp), grp in df.groupby(["dataset", "Species"], sort=False):
  576. t1 = grp["Species_start"].iloc[0]
  577. t2 = grp["Species_end"].iloc[0]
  578. draw_wedge(ax, t1, t2, r_species[0], r_species[1],
  579. species_colors.get(sp, "#e8e8e8"))
  580. mid_deg = (t1 + t2) / 2
  581. mid = np.deg2rad(mid_deg)
  582. r_t = (r_species[0] + r_species[1]) / 2
  583. norm = mid_deg % 360
  584. rot = mid_deg + 180 if 90 < norm < 270 else mid_deg
  585. ax.text(r_t * np.cos(mid), r_t * np.sin(mid), sp,
  586. ha="center", va="center", fontsize=9,
  587. rotation=rot, rotation_mode="anchor")
  588. # ── Ring 3: edge_type ─────────────────────────────────────────────────────
  589. for _, row in df.iterrows():
  590. t1, t2 = row["angle_start"], row["angle_end"]
  591. color = edge_colors.get(row["edge_type"], "#cccccc")
  592. draw_wedge(ax, t1, t2, r_edge[0], r_edge[1], color)
  593. # ── Outer count labels ────────────────────────────────────────────────────
  594. # IMPORTANT: multiple raw rows can share the same edge_type wedge (e.g. many
  595. # different relation sub-types all rolled into the single "relation" wedge),
  596. # and they all share the same angle_start/angle_end. Looping over raw rows
  597. # here would draw one label per underlying row, all stacked at the same
  598. # midpoint -> garbled overlapping digits. Aggregate to one label per wedge first.
  599. label_df = (df.groupby(["Species", "edge_type", "angle_start", "angle_end"], sort=False)
  600. ["count"].sum().reset_index())
  601. MIN_LABEL_DEG = 1.5
  602. skipped = []
  603. for _, row in label_df.iterrows():
  604. t1, t2 = row["angle_start"], row["angle_end"]
  605. if t2 - t1 <= 0:
  606. continue
  607. if (t2 - t1) < MIN_LABEL_DEG:
  608. skipped.append((row.get("Species", ""), row["edge_type"], row["count"]))
  609. continue
  610. mid_deg = (t1 + t2) / 2
  611. mid = np.deg2rad(mid_deg)
  612. norm = mid_deg % 360
  613. rot = mid_deg + 180 if 90 < norm < 270 else mid_deg
  614. ha = "right" if 90 < norm < 270 else "left"
  615. ax.text(r_label * np.cos(mid), r_label * np.sin(mid),
  616. format_count(row["count"]),
  617. ha=ha, va="center", fontsize=9, fontweight="bold", color="black",
  618. rotation=rot, rotation_mode="anchor",
  619. path_effects=[pe.withStroke(linewidth=2.5, foreground="white")])
  620. # list skipped (too-thin-to-label) values in the corner so the exact
  621. # numbers are still available, just not crammed onto the wedge
  622. if skipped:
  623. lines = [f"{sp} / {et}: {format_count(c)}" for sp, et, c in skipped]
  624. ax.text(-1.55, -1.55, "small slices (not labeled on ring):\n" + "\n".join(lines),
  625. ha="left", va="bottom", fontsize=6, color="dimgray")
  626. ax.set_xlim(-1.6, 1.6)
  627. ax.set_ylim(-1.6, 1.6)
  628. ax.axis("off")
  629. ax.set_title(title, fontsize=15, fontweight="bold", pad=20)
  630. # -----------------------------------------------------------------------------
  631. # 5. BUILD LOG-SCALED DATAFRAME (only)
  632. # -----------------------------------------------------------------------------
  633. df_log = assign_angles(df.copy(), ["dataset", "Species", "edge_type"])
  634. df_log["angle_start"] = df_log["edge_type_start"]
  635. df_log["angle_end"] = df_log["edge_type_end"]
  636. # -----------------------------------------------------------------------------
  637. # 6. PLOT — single log-scaled sunburst
  638. # -----------------------------------------------------------------------------
  639. fig, ax = plt.subplots(1, 1, figsize=(13, 13), subplot_kw={"aspect": "equal"})
  640. draw_sunburst(ax, df_log,
  641. "Edge-Type Counts — Log₁₀-scaled\n(wedge size ∝ log₁₀(count))")
  642. legend_handles = [
  643. mpatches.Patch(color=c, label=edge_legend_labels.get(k, k))
  644. for k, c in edge_colors.items()
  645. ]
  646. legend_handles += [
  647. mpatches.Patch(color=c, label=k)
  648. for k, c in dataset_colors.items()
  649. ]
  650. fig.legend(handles=legend_handles, loc="lower center", ncol=5,
  651. fontsize=11, bbox_to_anchor=(0.5, -0.03), frameon=False)
  652. fig.suptitle("Sunburst: Edge-Type Counts by Dataset & Species\n(AssociatedWith excluded)",
  653. fontsize=17, fontweight="bold", y=1.01)
  654. plt.tight_layout()
  655. plt.savefig("edge-type-sunburst_log_aging_biomedical.svg", format="svg", bbox_inches="tight")
  656. plt.show()
  657. df.head()
  658. # %% [markdown]
  659. # ## 1e
  660. # %%
  661. import matplotlib
  662. matplotlib.rcParams['svg.fonttype'] = 'none'
  663. matplotlib.rcParams['font.family'] = 'sans-serif'
  664. import matplotlib.pyplot as plt
  665. import matplotlib.patches as mpatches
  666. import matplotlib.lines as mlines
  667. import numpy as np
  668. import pandas as pd
  669. from matplotlib.patches import Circle, Ellipse
  670. # ── File paths ────────────────────────────────────────────────────────────────
  671. AGING_NODES_CSV = "/storage/Arushi/090526_EvoAge/kg_formation/species_wise_kg_building/aging_121_12M_Species_NodeType_unique_counts.csv"
  672. AGING_TRIPLES_CSV = "/storage/Arushi/090526_EvoAge/kg_formation/species_wise_kg_building/aging_121_12M_RelationType_AllSpecies.csv"
  673. BIO_NODES_CSV = "/storage/Arushi/090526_EvoAge/kg_formation/species_wise_kg_building/Species_Biomedical_121_12M_NodeType_unique_counts.csv"
  674. BIO_TRIPLES_CSV = "/storage/Arushi/090526_EvoAge/kg_formation/species_wise_kg_building/Species_Biomedical_121_12M_RelationType_unique_counts.csv"
  675. OUTPUT_SVG = "nodes_triples_dumbbell.svg"
  676. OUTPUT_PNG = "nodes_triples_dumbbell.png"
  677. # ── Colors ────────────────────────────────────────────────────────────────────
  678. AGING_COL = '#F08080' # salmon
  679. BIO_COL = '#80b4b9' # teal
  680. # ── Species order (left → right on x-axis, matching icon order) ───────────────
  681. # Column names as they appear in CSVs → display label
  682. SPECIES_MAP = {
  683. 'Human': 'Human',
  684. 'Mouse': 'Mouse',
  685. 'Zebrafish': 'Zebrafish',
  686. 'Celegans': 'C. elegans',
  687. 'Drosophila': 'Drosophila',
  688. 'Yeast': 'Yeast',
  689. }
  690. SPECIES_COLS = list(SPECIES_MAP.keys()) # CSV column names
  691. SPECIES_ORDER = list(SPECIES_MAP.keys()) # plot order
  692. # ── Load & extract TOTAL row ──────────────────────────────────────────────────
  693. def load_total_row(filepath, total_col_value):
  694. """
  695. Reads a CSV where the first column is a label column.
  696. Returns a dict {species_col: count} for the row whose
  697. first-column value matches `total_col_value`.
  698. """
  699. df = pd.read_csv(filepath)
  700. label_col = df.columns[0]
  701. row = df[df[label_col] == total_col_value].iloc[0]
  702. return {sp: int(row[sp]) if sp in df.columns else 0 for sp in SPECIES_COLS}
  703. aging_nodes = load_total_row(AGING_NODES_CSV, 'TOTAL')
  704. aging_triples = load_total_row(AGING_TRIPLES_CSV, 'Total Triples')
  705. bio_nodes = load_total_row(BIO_NODES_CSV, 'TOTAL')
  706. bio_triples = load_total_row(BIO_TRIPLES_CSV, 'Total Triples')
  707. # ── log10(count + 1) helper ───────────────────────────────────────────────────
  708. def log1(val):
  709. return np.log10(val + 1)
  710. # ── Figure ────────────────────────────────────────────────────────────────────
  711. x_pos = np.arange(len(SPECIES_ORDER), dtype=float)
  712. OFFSET = 0.15 # left stem = nodes, right stem = triples
  713. fig, ax = plt.subplots(figsize=(7.5, 6.2))
  714. fig.patch.set_facecolor('white')
  715. ax.set_facecolor('white')
  716. for i, sp in enumerate(SPECIES_ORDER):
  717. cx = x_pos[i]
  718. an = log1(aging_nodes[sp])
  719. bn = log1(bio_nodes[sp])
  720. at = log1(aging_triples[sp])
  721. bt = log1(bio_triples[sp])
  722. x_node = cx - OFFSET # left dumbbell → node counts
  723. x_triple = cx + OFFSET # right dumbbell → triple counts
  724. # vertical stems
  725. ax.plot([x_node, x_node], [an, bn], color='#333333', lw=1.3, zorder=1, solid_capstyle='round')
  726. ax.plot([x_triple, x_triple], [at, bt], color='#333333', lw=1.3, zorder=1, solid_capstyle='round')
  727. # node dumbbell: circles
  728. ax.scatter(x_node, an, marker='o', s=90, color=AGING_COL, zorder=4, linewidths=0.5, edgecolors='white')
  729. ax.scatter(x_node, bn, marker='o', s=90, color=BIO_COL, zorder=4, linewidths=0.5, edgecolors='white')
  730. # triple dumbbell: diamonds
  731. ax.scatter(x_triple, at, marker='D', s=70, color=AGING_COL, zorder=4, linewidths=0.5, edgecolors='white')
  732. ax.scatter(x_triple, bt, marker='D', s=70, color=BIO_COL, zorder=4, linewidths=0.5, edgecolors='white')
  733. # ── Axes formatting ───────────────────────────────────────────────────────────
  734. ax.set_ylabel(r'$\log_{10}\ (\mathrm{count}+1)$', fontsize=12)
  735. ax.set_xticks(x_pos)
  736. # Add species names as x-axis labels
  737. ax.set_xticklabels([SPECIES_MAP[sp] for sp in SPECIES_ORDER], fontsize=10, rotation=45, ha='right')
  738. ax.set_xlim(-0.6, len(SPECIES_ORDER) - 0.4)
  739. ax.set_ylim(-1.7, 10.5)
  740. ax.set_yticks(range(0, 11, 2))
  741. ax.tick_params(axis='y', labelsize=11)
  742. ax.spines['top'].set_visible(False)
  743. ax.spines['right'].set_visible(False)
  744. ax.spines['bottom'].set_linewidth(1.2)
  745. ax.spines['left'].set_linewidth(1.2)
  746. ax.set_title('Distribution (Nodes, Triples)', fontsize=13, pad=10)
  747. # ── Legend ────────────────────────────────────────────────────────────────────
  748. aging_patch = mpatches.Patch(color=AGING_COL, label='Aging')
  749. bio_patch = mpatches.Patch(color=BIO_COL, label='Biomedical')
  750. node_marker = mlines.Line2D([], [], color='grey', marker='o', linestyle='None', markersize=8, label='Node count')
  751. tri_marker = mlines.Line2D([], [], color='grey', marker='D', linestyle='None', markersize=7, label='Triple count')
  752. ax.legend(
  753. handles=[aging_patch, node_marker, bio_patch, tri_marker],
  754. loc='upper right', fontsize=10, frameon=False, ncol=2,
  755. handlelength=1.2, columnspacing=0.8, labelspacing=0.4,
  756. )
  757. # ── Save ──────────────────────────────────────────────────────────────────────
  758. plt.tight_layout()
  759. plt.savefig(OUTPUT_SVG, format='svg', bbox_inches='tight')
  760. plt.savefig(OUTPUT_PNG, dpi=200, bbox_inches='tight')
  761. print(f"Saved: {OUTPUT_SVG}, {OUTPUT_PNG}")
  762. # %% [markdown]
  763. # ## 1f
  764. # %%
  765. # /storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/Graph_features/Output_csv/Aging_121_12M/Aging_121_12M_degree.csv
  766. # /storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/Graph_features/Output_csv/EvoAge_121_12M/EvoAge_121_12M_degree.csv
  767. import pandas as pd
  768. import numpy as np
  769. import matplotlib.pyplot as plt
  770. from matplotlib.ticker import LogLocator, FuncFormatter
  771. # -----------------------------------------------------------------------------
  772. # 1. LOAD DATA
  773. # -----------------------------------------------------------------------------
  774. aging_path = "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/Graph_features/Output_csv/Aging_121_12M/Aging_121_12M_degree.csv"
  775. evoage_path = "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/Graph_features/Output_csv/EvoAge_121_12M/EvoAge_121_12M_degree.csv"
  776. aging = pd.read_csv(aging_path)
  777. evoage = pd.read_csv(evoage_path)
  778. # -----------------------------------------------------------------------------
  779. # 2. BUILD NORMALIZED DEGREE DISTRIBUTION: P(k) = fraction of nodes with degree k
  780. # -----------------------------------------------------------------------------
  781. def to_distribution(df, model_name):
  782. N = len(df)
  783. dist = (
  784. df["total_degree"]
  785. .value_counts()
  786. .rename_axis("degree")
  787. .reset_index(name="num_nodes")
  788. )
  789. dist["P_k"] = dist["num_nodes"] / N # normalize by total node count
  790. dist["Model"] = model_name
  791. dist["N"] = N
  792. return dist
  793. aging_dist = to_distribution(aging, "Aging")
  794. evoage_dist = to_distribution(evoage, "EvoAge")
  795. N_aging = len(aging)
  796. N_evoage = len(evoage)
  797. data_long = pd.concat([evoage_dist, aging_dist], ignore_index=True)
  798. data_long = data_long[(data_long["degree"] > 0) & (data_long["P_k"] > 0)]
  799. # -----------------------------------------------------------------------------
  800. # 3. PLOT (log-log scatter)
  801. # -----------------------------------------------------------------------------
  802. colors = {"EvoAge": "#2c7189", "Aging": "#b26f77"}
  803. labels = {"EvoAge": f"EvoAge (N = {N_evoage:,})",
  804. "Aging": f"Aging (N = {N_aging:,})"}
  805. fig, ax = plt.subplots(figsize=(10, 8))
  806. for model in ["EvoAge", "Aging"]:
  807. subset = data_long[data_long["Model"] == model]
  808. ax.scatter(
  809. subset["degree"], subset["P_k"],
  810. color=colors[model], s=45, alpha=0.7,
  811. edgecolor="none", label=labels[model],
  812. )
  813. ax.set_xscale("log")
  814. ax.set_yscale("log")
  815. ax.xaxis.set_major_locator(LogLocator(base=10, numticks=15))
  816. ax.yaxis.set_major_locator(LogLocator(base=10, numticks=15))
  817. ax.xaxis.set_major_formatter(FuncFormatter(lambda x, _: f"$10^{{{int(np.log10(x))}}}$"))
  818. ax.yaxis.set_major_formatter(FuncFormatter(lambda x, _: f"$10^{{{int(np.log10(x))}}}$"))
  819. # Auto-fit limits to actual data
  820. max_deg = data_long["degree"].max()
  821. ax.set_xlim(1, 10 ** np.ceil(np.log10(max_deg)))
  822. min_pk = data_long["P_k"].min()
  823. ax.set_ylim(10 ** np.floor(np.log10(min_pk)), 1)
  824. ax.set_title("Comparative Log-Log Plot of Degree Distributions",
  825. fontsize=20, fontweight="bold", pad=15)
  826. ax.text(0.5, 1.02, "Normalized scatter representation on log-log scale",
  827. transform=ax.transAxes, ha="center", fontsize=12, color="gray")
  828. ax.set_xlabel("Node Degree, $k$ (Log Scale)", fontsize=16)
  829. ax.set_ylabel("Fraction of Nodes, $P(k)$ (Log Scale)", fontsize=16)
  830. ax.spines["top"].set_visible(False)
  831. ax.spines["right"].set_visible(False)
  832. ax.spines["left"].set_color("black")
  833. ax.spines["bottom"].set_color("black")
  834. ax.tick_params(axis="both", which="major", length=6, width=1, color="black", labelsize=13)
  835. ax.tick_params(axis="both", which="minor", length=0)
  836. ax.grid(False, which="major")
  837. ax.grid(False, which="minor")
  838. ax.legend(loc="lower center", bbox_to_anchor=(0.5, -0.18),
  839. ncol=2, frameon=False, fontsize=14, markerscale=1.5)
  840. plt.tight_layout()
  841. # -----------------------------------------------------------------------------
  842. # 4. SAVE AS PDF
  843. # -----------------------------------------------------------------------------
  844. plt.savefig("FIG1/Fig1_h_main.pdf", format="pdf", bbox_inches="tight")
  845. plt.savefig("FIG1/fig1h-degree.png", format="png", bbox_inches="tight")
  846. plt.show()
  847. # %% [markdown]
  848. # ## 1g
  849. # %%
  850. import pandas as pd
  851. import numpy as np
  852. import matplotlib.pyplot as plt
  853. # -----------------------------------------------------------------------------
  854. # 1. LOAD DATA
  855. # -----------------------------------------------------------------------------
  856. aging_path = "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/Graph_features/Output_csv/Aging_121_12M/Aging_121_12M_louvain_partitions.csv"
  857. evoage_path = "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/Graph_features/Output_csv/Biomedical_121_12M/Biomedical_121_12M_louvain_partitions.csv"
  858. def read_louvain_csv(path):
  859. df = pd.read_csv(path)
  860. df.columns = df.columns.str.lower()
  861. if not {"vertex", "partition"}.issubset(df.columns):
  862. raise ValueError(f"Expected 'vertex' and 'partition' columns in {path}, got {list(df.columns)}")
  863. return df[["vertex", "partition"]]
  864. aging_raw = read_louvain_csv(aging_path)
  865. evoage_raw = read_louvain_csv(evoage_path)
  866. # -----------------------------------------------------------------------------
  867. # 2. COMMUNITY SIZE TABLES
  868. # -----------------------------------------------------------------------------
  869. def community_sizes(df):
  870. return (
  871. df.groupby("partition")
  872. .size()
  873. .reset_index(name="size")
  874. .sort_values("size", ascending=False)
  875. .reset_index(drop=True)
  876. )
  877. aging_sizes = community_sizes(aging_raw)
  878. evoage_sizes = community_sizes(evoage_raw)
  879. N_aging = len(aging_raw)
  880. N_evoage = len(evoage_raw)
  881. # -----------------------------------------------------------------------------
  882. # 3. QUICK STATS (mirrors R's cat block)
  883. # -----------------------------------------------------------------------------
  884. def entropy(sizes):
  885. p = sizes / sizes.sum()
  886. p = p[p > 0]
  887. return float(-np.sum(p * np.log(p)))
  888. print("\n=== Summary ===")
  889. print(f"Aging | Nodes: {N_aging:,} | Communities: {len(aging_sizes):,} | "
  890. f"Largest: {100*aging_sizes['size'].max()/N_aging:.2f}% | "
  891. f"Entropy: {entropy(aging_sizes['size'].values):.3f}")
  892. print(f"Biomedical | Nodes: {N_evoage:,} | Communities: {len(evoage_sizes):,} | "
  893. f"Largest: {100*evoage_sizes['size'].max()/N_evoage:.2f}% | "
  894. f"Entropy: {entropy(evoage_sizes['size'].values):.3f}\n")
  895. # -----------------------------------------------------------------------------
  896. # 4. TOP-K COMMUNITIES
  897. # -----------------------------------------------------------------------------
  898. K = 15
  899. top_aging = aging_sizes.head(K).copy()
  900. top_aging["rank"] = np.arange(1, len(top_aging) + 1)
  901. top_aging["pct"] = 100 * top_aging["size"] / N_aging
  902. top_aging["graph"] = "Aging"
  903. top_evoage = evoage_sizes.head(K).copy()
  904. top_evoage["rank"] = np.arange(1, len(top_evoage) + 1)
  905. top_evoage["pct"] = 100 * top_evoage["size"] / N_evoage
  906. top_evoage["graph"] = "Biomedical"
  907. top_df = pd.concat([top_aging, top_evoage], ignore_index=True)
  908. # -----------------------------------------------------------------------------
  909. # 5. PLOT — grouped bar chart
  910. # -----------------------------------------------------------------------------
  911. colors = {"Aging": "#eeb4b4", "Biomedical": "#7ebcdb"}
  912. fig, ax = plt.subplots(figsize=(10, 6))
  913. ranks = np.arange(1, K + 1)
  914. width = 0.4
  915. aging_vals = top_df[top_df["graph"] == "Aging"].set_index("rank").reindex(ranks)["pct"].fillna(0)
  916. evoage_vals = top_df[top_df["graph"] == "Biomedical"].set_index("rank").reindex(ranks)["pct"].fillna(0)
  917. ax.bar(ranks - width/2, aging_vals, width=width, color=colors["Aging"], label="Aging",
  918. edgecolor="none")
  919. ax.bar(ranks + width/2, evoage_vals, width=width, color=colors["Biomedical"], label="Biomedical",
  920. edgecolor="none")
  921. ax.set_xticks(ranks)
  922. ax.set_xticklabels(ranks, fontsize=11)
  923. ax.set_title(f"Top-{K} Community Sizes (as % of nodes)",
  924. fontsize=15, fontweight="bold", pad=10)
  925. ax.text(0.5, 1.02,
  926. f"Top-{K} communities by size, shown as % of total nodes. "
  927. "Larger bars indicate more concentration in a few communities.",
  928. transform=ax.transAxes, ha="center", fontsize=10, color="gray")
  929. ax.set_xlabel("Community rank (by size, 1 = largest)", fontsize=13)
  930. ax.set_ylabel("Share of nodes in community (%)", fontsize=13)
  931. # theme_classic-ish
  932. ax.spines["top"].set_visible(False)
  933. ax.spines["right"].set_visible(False)
  934. ax.spines["left"].set_linewidth(0.7)
  935. ax.spines["bottom"].set_linewidth(0.7)
  936. ax.tick_params(axis="both", which="major", length=5, width=0.7, color="black", labelsize=11)
  937. ax.grid(False)
  938. ax.legend(loc="upper right", frameon=False, fontsize=12, ncol=2)
  939. plt.tight_layout()
  940. # -----------------------------------------------------------------------------
  941. # 6. SAVE AS PDF
  942. # -----------------------------------------------------------------------------
  943. plt.savefig("FIG1/aging-biomedical-commmunity.pdf", format="pdf", bbox_inches="tight")
  944. plt.show()
  945. # %% [markdown]
  946. # ## 1h
  947. # %%
  948. import pandas as pd
  949. import numpy as np
  950. import matplotlib.pyplot as plt
  951. from matplotlib.ticker import LogLocator, FuncFormatter
  952. # -----------------------------------------------------------------------------
  953. # 1. LOAD DATA
  954. # -----------------------------------------------------------------------------
  955. aging_path = "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/Graph_features/Output_csv/Aging_121_12M/Aging_121_12M_betweenness_centrality_sorted.csv"
  956. evoage_path = "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/Graph_features/Output_csv/Biomedical_121_12M/Biomedical_121_12M_betweenness_centrality_sorted.csv"
  957. df_aging = pd.read_csv(aging_path)
  958. df_evoage = pd.read_csv(evoage_path)
  959. # -----------------------------------------------------------------------------
  960. # 2. CONFIG
  961. # -----------------------------------------------------------------------------
  962. TARGET_ROWS = 2_000_000
  963. NBINS_PDF = 120
  964. SEED = 42
  965. # -----------------------------------------------------------------------------
  966. # 3. HELPERS
  967. # -----------------------------------------------------------------------------
  968. def positive_vec(x, min_positive=1e-12):
  969. """Keep only finite, positive values (required for log-scale)."""
  970. x = pd.to_numeric(x, errors="coerce").to_numpy()
  971. x = x[np.isfinite(x)]
  972. x = x[x > 0]
  973. return x if len(x) else np.array([min_positive])
  974. def log_binned_pdf(arr, nbins=120):
  975. """Width-corrected PDF using log-spaced bins."""
  976. lo, hi = arr.min(), arr.max()
  977. if hi <= lo:
  978. return np.array([hi]), np.array([1.0])
  979. bins = np.logspace(np.log10(lo), np.log10(hi), nbins + 1)
  980. counts, edges = np.histogram(arr, bins=bins)
  981. widths = np.diff(edges)
  982. total = counts.sum()
  983. pdf = np.where(widths > 0, (counts / total) / widths, 0)
  984. centers = np.sqrt(edges[:-1] * edges[1:]) # geometric mean of bin edges
  985. # keep only bins with non-zero density (log axis can't plot 0)
  986. mask = pdf > 0
  987. return centers[mask], pdf[mask]
  988. # -----------------------------------------------------------------------------
  989. # 4. PREP
  990. # -----------------------------------------------------------------------------
  991. # Subsample EvoAge if too large (matches R's TARGET_ROWS = 2e6 on df2)
  992. if len(df_evoage) > TARGET_ROWS:
  993. df_evoage = df_evoage.sample(n=TARGET_ROWS, random_state=SEED)
  994. a_aging = positive_vec(df_aging["betweenness_centrality"])
  995. a_evoage = positive_vec(df_evoage["betweenness_centrality"])
  996. x_aging, y_aging = log_binned_pdf(a_aging, NBINS_PDF)
  997. x_evoage, y_evoage = log_binned_pdf(a_evoage, NBINS_PDF)
  998. # -----------------------------------------------------------------------------
  999. # 5. PLOT
  1000. # -----------------------------------------------------------------------------
  1001. colors = {"Aging": "#eeb4b4", "Biomedical": "#7ebcdb"}
  1002. fig, ax = plt.subplots(figsize=(8, 7))
  1003. ax.plot(x_evoage, y_evoage, color=colors["Biomedical"], linewidth=1.4, label="Biomedical")
  1004. ax.plot(x_aging, y_aging, color=colors["Aging"], linewidth=1.4, label="Aging")
  1005. ax.set_xscale("log")
  1006. ax.set_yscale("log")
  1007. # 10^n tick formatting
  1008. def sci_fmt(x, _):
  1009. return f"$10^{{{int(np.log10(x))}}}$"
  1010. ax.xaxis.set_major_locator(LogLocator(base=10, numticks=20))
  1011. ax.yaxis.set_major_locator(LogLocator(base=10, numticks=20))
  1012. ax.xaxis.set_major_formatter(FuncFormatter(sci_fmt))
  1013. ax.yaxis.set_major_formatter(FuncFormatter(sci_fmt))
  1014. # Labels
  1015. ax.set_title("connectivity", fontsize=16, pad=12)
  1016. ax.set_xlabel("betweenness centrality", fontsize=14)
  1017. ax.set_ylabel("density", fontsize=14)
  1018. # Rotate x tick labels a bit (matches your reference image)
  1019. plt.setp(ax.get_xticklabels(), rotation=45, ha="right")
  1020. # Clean minimal look
  1021. ax.spines["top"].set_visible(False)
  1022. ax.spines["right"].set_visible(False)
  1023. ax.tick_params(axis="both", which="major", length=5, width=1, color="black", labelsize=11)
  1024. ax.tick_params(axis="both", which="minor", length=0)
  1025. ax.grid(False)
  1026. # Legend top-center inside plot
  1027. ax.legend(loc="upper center", bbox_to_anchor=(0.5, 0.98),
  1028. ncol=2, frameon=False, fontsize=13)
  1029. plt.tight_layout()
  1030. # -----------------------------------------------------------------------------
  1031. # 6. SAVE AS PDF
  1032. # -----------------------------------------------------------------------------
  1033. plt.savefig("FIG1/aging-biomedical-centrality.pdf", format="pdf", bbox_inches="tight")
  1034. plt.show()
  1035. # %% [markdown]
  1036. # ## 1i
  1037. # %%
  1038. import pandas as pd
  1039. import numpy as np
  1040. import matplotlib.pyplot as plt
  1041. from matplotlib.ticker import LogLocator, FuncFormatter
  1042. # -----------------------------------------------------------------------------
  1043. # 1. LOAD DATA
  1044. # -----------------------------------------------------------------------------
  1045. aging_path = "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/Graph_features/Output_csv/Aging_121_12M_vertex_counts.csv"
  1046. evoage_path = "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/Graph_features/Output_csv/Biomedical_121_12M/Biomedical_121_12M_vertex_counts.csv"
  1047. def read_tri_csv(path):
  1048. df = pd.read_csv(path)
  1049. df.columns = df.columns.str.lower()
  1050. if not {"vertex", "counts"}.issubset(df.columns):
  1051. raise ValueError(f"Expected 'vertex' and 'counts' columns in {path}, got {list(df.columns)}")
  1052. return df[["vertex", "counts"]].apply(pd.to_numeric, errors="coerce")
  1053. aging_raw = read_tri_csv(aging_path)
  1054. evoage_raw = read_tri_csv(evoage_path)
  1055. N_aging = len(aging_raw)
  1056. N_evoage = len(evoage_raw)
  1057. # -----------------------------------------------------------------------------
  1058. # 2. QUICK STATS
  1059. # -----------------------------------------------------------------------------
  1060. aging_zero = (aging_raw["counts"] <= 0).mean()
  1061. evoage_zero = (evoage_raw["counts"] <= 0).mean()
  1062. print("\n=== Triangle Count Summary ===")
  1063. print(f"Aging | Nodes: {N_aging:,} | zero-triangle nodes: {100*aging_zero:.2f}%")
  1064. print(f"Biomedical | Nodes: {N_evoage:,} | zero-triangle nodes: {100*evoage_zero:.2f}%\n")
  1065. # Keep positive counts for log-scale
  1066. pos_aging = aging_raw.loc[aging_raw["counts"] > 0, "counts"].to_numpy()
  1067. pos_evoage = evoage_raw.loc[evoage_raw["counts"] > 0, "counts"].to_numpy()
  1068. # -----------------------------------------------------------------------------
  1069. # 3. LOG-BINNED PDF (matches R's hist(..., include.lowest=TRUE, right=FALSE))
  1070. # -----------------------------------------------------------------------------
  1071. def log_binned_pdf(arr, nbins=120):
  1072. arr = arr[np.isfinite(arr) & (arr > 0)]
  1073. if len(arr) == 0:
  1074. return np.array([1.0]), np.array([1.0])
  1075. lo, hi = arr.min(), arr.max()
  1076. if hi <= lo:
  1077. return np.array([hi]), np.array([1.0])
  1078. bins = np.logspace(np.log10(lo), np.log10(hi), nbins + 1)
  1079. # Replicate R's right=FALSE, include.lowest=TRUE:
  1080. # left-closed, right-open intervals [a,b); include max value in last bin
  1081. idx = np.digitize(arr, bins, right=False)
  1082. idx[arr == bins[-1]] = nbins # include the maximum value in the last bin
  1083. counts = np.bincount(idx, minlength=nbins + 2)[1:nbins + 1]
  1084. widths = np.diff(bins)
  1085. total = counts.sum()
  1086. pdf = np.where(widths > 0, (counts / total) / widths, 0)
  1087. centers = np.sqrt(bins[:-1] * bins[1:])
  1088. mask = pdf > 0
  1089. return centers[mask], pdf[mask]
  1090. x_aging, y_aging = log_binned_pdf(pos_aging, nbins=120)
  1091. x_evoage, y_evoage = log_binned_pdf(pos_evoage, nbins=120)
  1092. # -----------------------------------------------------------------------------
  1093. # 4. PLOT
  1094. # -----------------------------------------------------------------------------
  1095. colors = {"Aging": "#eeb4b4", "Biomedical": "#7ebcdb"}
  1096. fig, ax = plt.subplots(figsize=(8, 7))
  1097. # Aging as dashed line (matches your reference image)
  1098. ax.plot(x_aging, y_aging, color=colors["Aging"],
  1099. linewidth=1.4, linestyle="--", label="Aging")
  1100. ax.plot(x_evoage, y_evoage, color=colors["Biomedical"],
  1101. linewidth=1.4, linestyle="-", label="Biomedical")
  1102. ax.set_xscale("log")
  1103. ax.set_yscale("log")
  1104. def sci_fmt(x, _):
  1105. return f"$10^{{{int(np.log10(x))}}}$"
  1106. ax.xaxis.set_major_locator(LogLocator(base=10, numticks=20))
  1107. ax.yaxis.set_major_locator(LogLocator(base=10, numticks=20))
  1108. ax.xaxis.set_major_formatter(FuncFormatter(sci_fmt))
  1109. ax.yaxis.set_major_formatter(FuncFormatter(sci_fmt))
  1110. ax.set_title("clustering", fontsize=16, pad=12)
  1111. ax.set_xlabel("triangles per node", fontsize=14)
  1112. ax.set_ylabel("density", fontsize=14)
  1113. ax.spines["top"].set_visible(False)
  1114. ax.spines["right"].set_visible(False)
  1115. ax.tick_params(axis="both", which="major", length=5, width=1, color="black", labelsize=11)
  1116. ax.tick_params(axis="both", which="minor", length=0)
  1117. ax.grid(False)
  1118. ax.legend(loc="upper center", bbox_to_anchor=(0.5, 0.98),
  1119. ncol=2, frameon=False, fontsize=13)
  1120. plt.tight_layout()
  1121. plt.show()
  1122. # -----------------------------------------------------------------------------
  1123. # 5. SAVE AS PDF
  1124. # -----------------------------------------------------------------------------
  1125. plt.savefig("FIG1/aging-biomedical-clustering.pdf", format="pdf", bbox_inches="tight")
  1126. # %% [markdown]
  1127. # ## Main Figure 2
  1128. # %% [markdown]
  1129. # ## 2b
  1130. # %%
  1131. """
  1132. 3-set Venn across both KG variants (1:1 vs 121_12M, sharing the same test set),
  1133. for all three KG families — Aging, Biomedical, EvoAge.
  1134. """
  1135. import gc
  1136. import numpy as np
  1137. import torch
  1138. import matplotlib.pyplot as plt
  1139. from matplotlib_venn import venn3
  1140. # ============================================================
  1141. # CONFIG
  1142. # ============================================================
  1143. FAMILIES = {
  1144. "Aging": {
  1145. "full_1to1": "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/building_aging_kg_new/Store_House/Aging_specific_1_to_1_KG.pt",
  1146. "test": "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/building_aging_kg_new/Store_House/Aging_specific_1to1_KG_test_10_shared.pt",
  1147. "full_12m": "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/building_aging_kg_new/Store_House/Aging_specific_121_12M_KG.pt",
  1148. },
  1149. "Biomedical": {
  1150. "full_1to1": "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/building_biomedical_kg_new/Store_House/Biomedical_1_to_1_KG.pt",
  1151. "test": "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/building_biomedical_kg_new/Store_House/Biomedical_1to1_KG_test_10_shared.pt",
  1152. "full_12m": "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/building_biomedical_kg_new/Store_House/Biomedical_121_12M_KG.pt",
  1153. },
  1154. "EvoAge": {
  1155. "full_1to1": "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/building_evoage_kg_new/Store_House/EvoAge_1_to_1_KG.pt",
  1156. "test": "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/building_evoage_kg_new/Store_House/EvoAge_1to1_KG_test_10_shared.pt",
  1157. "full_12m": "/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/building_evoage_kg_new/Store_House/EvoAge_121_12M_to_many_KG.pt",
  1158. },
  1159. }
  1160. ASSUME_UNIQUE_INPUT = False # True -> sort only, skip dedupe (if files already unique)
  1161. UNWEIGHTED = True # True -> readable fixed layout; False -> area-proportional
  1162. USE_GPU = False # True -> CuPy path (needs cupy + enough VRAM per set)
  1163. xp = np
  1164. if USE_GPU:
  1165. import cupy as cp
  1166. xp = cp
  1167. # ============================================================
  1168. # LOADER -> (N,3) int64 on CPU
  1169. # ============================================================
  1170. def load_triples_pt(path):
  1171. obj = torch.load(path, map_location="cpu", weights_only=False)
  1172. if isinstance(obj, torch.Tensor) and obj.ndim == 2 and obj.shape[1] == 3:
  1173. return obj.numpy().astype(np.int64, copy=False)
  1174. if isinstance(obj, dict):
  1175. for key in ("triples", "edge_triples", "all_triples"):
  1176. if key in obj:
  1177. return np.asarray(obj[key], dtype=np.int64)
  1178. if "edge_index" in obj and "edge_type" in obj:
  1179. h, t = np.asarray(obj["edge_index"]); r = np.asarray(obj["edge_type"])
  1180. return np.stack([h, r, t], axis=1).astype(np.int64)
  1181. if hasattr(obj, "edge_index") and hasattr(obj, "edge_type"):
  1182. h, t = obj.edge_index.numpy(); r = obj.edge_type.numpy()
  1183. return np.stack([h, r, t], axis=1).astype(np.int64)
  1184. raise ValueError(f"Unrecognized structure in {path}: type={type(obj)}")
  1185. # ============================================================
  1186. # helpers
  1187. # ============================================================
  1188. def unique_keys(triples, r_max, t_max):
  1189. h = xp.asarray(triples[:, 0]); r = xp.asarray(triples[:, 1]); t = xp.asarray(triples[:, 2])
  1190. keys = (h * r_max + r) * t_max + t # bijective packing
  1191. return xp.sort(keys) if ASSUME_UNIQUE_INPUT else xp.unique(keys)
  1192. def isin_mask(query, ref):
  1193. """Which elements of sorted query are present in sorted ref."""
  1194. if ref.size == 0:
  1195. return xp.zeros(query.size, dtype=bool)
  1196. idx = xp.clip(xp.searchsorted(ref, query), 0, ref.size - 1)
  1197. return ref[idx] == query
  1198. def count_in(query, ref): # pass the SMALLER array as query
  1199. return int(isin_mask(query, ref).sum())
  1200. # version-robust unweighted layout
  1201. def make_venn3(subsets, labels, ax, unweighted=True):
  1202. if not unweighted:
  1203. return venn3(subsets=subsets, set_labels=labels, ax=ax)
  1204. try:
  1205. # newer matplotlib_venn (>=1.0): fixed_subset_sizes, no normalize_to
  1206. from matplotlib_venn.layout.venn3 import DefaultLayoutAlgorithm
  1207. layout = DefaultLayoutAlgorithm(fixed_subset_sizes=(1, 1, 1, 1, 1, 1, 1))
  1208. return venn3(subsets=subsets, set_labels=labels, ax=ax, layout_algorithm=layout)
  1209. except ImportError:
  1210. # older matplotlib_venn: the deprecated helper still works
  1211. from matplotlib_venn import venn3_unweighted
  1212. return venn3_unweighted(subsets=subsets, set_labels=labels, ax=ax)
  1213. # ============================================================
  1214. # per-family Venn
  1215. # ============================================================
  1216. def venn_for_family(family, paths):
  1217. raw = {
  1218. "Full 1:1 KG": load_triples_pt(paths["full_1to1"]),
  1219. "Test (shared)": load_triples_pt(paths["test"]),
  1220. "Full 121_12M KG": load_triples_pt(paths["full_12m"]),
  1221. }
  1222. # shared packing bases within THIS family so keys are comparable across its three sets
  1223. r_max = max(int(a[:, 1].max()) for a in raw.values()) + 1
  1224. t_max = max(int(a[:, 2].max()) for a in raw.values()) + 1
  1225. h_max = max(int(a[:, 0].max()) for a in raw.values())
  1226. assert h_max * r_max * t_max + (r_max - 1) * t_max + (t_max - 1) < 2**63 - 1, \
  1227. f"[{family}] packing overflows int64 — switch to a row-void view for keys"
  1228. keysets = {}
  1229. for name in list(raw):
  1230. keysets[name] = unique_keys(raw[name], r_max, t_max)
  1231. del raw[name]; gc.collect() # free each large raw array immediately
  1232. A, B, C = keysets["Full 1:1 KG"], keysets["Test (shared)"], keysets["Full 121_12M KG"]
  1233. a, b, c = A.size, B.size, C.size
  1234. # intersections — always query the smaller array
  1235. inBA = isin_mask(B, A); ab = int(inBA.sum())
  1236. AB = B[inBA] # A∩B keys (small)
  1237. bc = count_in(B, C)
  1238. ac = count_in(A, C) if a <= c else count_in(C, A)
  1239. abc = count_in(AB, C)
  1240. # 7 disjoint regions
  1241. ABC = abc
  1242. AB_only = ab - abc
  1243. AC_only = ac - abc
  1244. BC_only = bc - abc
  1245. A_only = a - ab - ac + abc
  1246. B_only = b - ab - bc + abc
  1247. C_only = c - ac - bc + abc
  1248. print(f"\n=== {family} ===")
  1249. for name, s in [("Full 1:1 KG", a), ("Test (shared)", b), ("Full 121_12M KG", c)]:
  1250. print(f" {name:18s}: {s:,} unique triples")
  1251. print(f" A∩B={ab:,} A∩C={ac:,} B∩C={bc:,} A∩B∩C={abc:,}")
  1252. # matplotlib_venn subset order: (Abc, aBc, ABc, abC, AbC, aBC, ABC)
  1253. subsets = (A_only, B_only, AB_only, C_only, AC_only, BC_only, ABC)
  1254. labels = ("Full 1:1 KG", "Test (shared)", "Full 121_12M KG")
  1255. fig, ax = plt.subplots(figsize=(9, 9))
  1256. make_venn3(subsets, labels, ax, unweighted=UNWEIGHTED)
  1257. ax.set_title(f"{family}: triple-set provenance across KG variants (real counts)",
  1258. fontsize=15, fontweight="bold")
  1259. fig.tight_layout()
  1260. out = f"venn_3set_{family}.svg"
  1261. fig.savefig(out, bbox_inches="tight")
  1262. print(f" saved -> {out}")
  1263. return fig
  1264. # %%
  1265. family = "Aging"
  1266. fig_aging = venn_for_family(family, FAMILIES[family])
  1267. plt.show()
  1268. # %%
  1269. family = "Biomedical"
  1270. fig_biomedical = venn_for_family(family, FAMILIES[family])
  1271. plt.show()
  1272. # %%
  1273. family = "EvoAge"
  1274. fig_evoage = venn_for_family(family, FAMILIES[family])
  1275. plt.show()
  1276. # %% [markdown]
  1277. # ## 2c
  1278. # %%
  1279. import pandas as pd
  1280. import numpy as np
  1281. import matplotlib.pyplot as plt
  1282. import matplotlib.gridspec as gridspec
  1283. from matplotlib.colors import LinearSegmentedColormap
  1284. # ----------------------------------------------------------------------
  1285. # 1. Load data
  1286. # ----------------------------------------------------------------------
  1287. data = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/training_3/Training_Stats/Trained_model_performance_Fig2b.csv')
  1288. # ----------------------------------------------------------------------
  1289. # 2. Clean / normalize model + graph names
  1290. # ----------------------------------------------------------------------
  1291. model_map = {
  1292. 'complex': 'ComplEx',
  1293. 'complexe': 'ComplEx',
  1294. 'simple': 'SimplE',
  1295. 'simpple': 'SimplE',
  1296. 'dismult': 'DisMult',
  1297. 'distmult': 'DisMult',
  1298. 'rescal': 'RESCAL',
  1299. 'rotate': 'RotatE',
  1300. 'transe': 'TransE',
  1301. }
  1302. data['Model'] = data['Model Type'].str.strip().str.lower().map(model_map)
  1303. graph_map = {
  1304. 'EvoAge_121_12M': 'EvoAge',
  1305. 'Aging_121_12M': 'Aging',
  1306. }
  1307. data['Graph'] = data['Graph'].map(graph_map)
  1308. data['metric_type'] = data['metric_type'].str.strip().str.title() # Validation / Testing
  1309. model_order = ['TransE', 'ComplEx', 'RotatE', 'SimplE', 'DisMult', 'RESCAL']
  1310. graph_order = ['Aging', 'EvoAge']
  1311. # ----------------------------------------------------------------------
  1312. # 3. FIGURE 2b — heatmap grid (EvoAge only, MRR + Hit@k)
  1313. # ----------------------------------------------------------------------
  1314. row_order = [('Validation', 'MRR'), ('Validation', 'Hit@1'), ('Validation', 'Hit@3'), ('Validation', 'Hit@10'),
  1315. ('Testing', 'MRR'), ('Testing', 'Hit@1'), ('Testing', 'Hit@3'), ('Testing', 'Hit@10')]
  1316. row_labels = ['MRR', 'hit@1', 'hit@3', 'hit@10', 'MRR', 'hit@1', 'hit@3', 'hit@10']
  1317. n_rows = len(row_order)
  1318. fig = plt.figure(figsize=(10, 4.5))
  1319. gs = gridspec.GridSpec(1, len(model_order) + 1,
  1320. width_ratios=[1] * len(model_order) + [0.12],
  1321. wspace=0.15)
  1322. vmin, vmax = 0.0, 1.0
  1323. # cmap = LinearSegmentedColormap.from_list('original_paper', ['#FCFDDD', '#A9DDD1', '#5EBBA6'])
  1324. # # ----- COLOR GRADIENT: edit these 3 hex codes -----
  1325. # COLOR_LOW = '#FCFDDD' # lowest values (0.0)
  1326. # COLOR_MID = '#A9DDD1' # midpoint (0.5)
  1327. # COLOR_HIGH = '#5EBBA6' # highest values (1.0)
  1328. # COLOR_LOW = '#F7FBFF'
  1329. # COLOR_MID = '#6BAED6'
  1330. # COLOR_HIGH = '#08306B'
  1331. # # cmap = LinearSegmentedColormap.from_list(
  1332. # # 'custom3', [COLOR_LOW, COLOR_MID, COLOR_HIGH]
  1333. # # )
  1334. # cmap = LinearSegmentedColormap.from_list(
  1335. # 'custom3',
  1336. # [(0.0, COLOR_LOW), (0.85, COLOR_MID), (1.0, COLOR_HIGH)] # change 0.5 to shift the middle color
  1337. # )
  1338. # ----- COLOR GRADIENT: 5 stops -----
  1339. C1 = '#F5D49B' # sand
  1340. C2 = '#F0AAAE' # pink
  1341. C3 = '#D98CC8' # mauve
  1342. C4 = '#C063E8' # violet
  1343. C5 = '#9B51E0' # purple
  1344. cmap = LinearSegmentedColormap.from_list('custom5', [C1, C2, C3, C4, C5])
  1345. # # ----- COLOR GRADIENT: 5 stops, edit these hex codes -----
  1346. # C1 = '#FFFFE5' # 0.00 lowest
  1347. # C2 = '#D9F0A3' # 0.25
  1348. # C3 = '#78C679' # 0.50
  1349. # C4 = '#238443' # 0.75
  1350. # C5 = '#004529' # 1.00 highest
  1351. # cmap = LinearSegmentedColormap.from_list('custom5', [C1, C2, C3, C4, C5])
  1352. # ----- COLOR GRADIENT: 5 stops, pastel purple end -----
  1353. C1 = '#F7DFB4' # soft sand
  1354. C2 = '#F2BCBE' # soft pink
  1355. C3 = '#DBA6D2' # soft mauve
  1356. C4 = '#C79BE0' # pastel violet
  1357. C5 = '#AE8CD9' # muted purple
  1358. cmap = LinearSegmentedColormap.from_list('custom5', [
  1359. (0.00, C1),
  1360. (0.20, C2),
  1361. (0.80, C3),
  1362. (0.90, C4),
  1363. (1.00, C5),
  1364. ])
  1365. im = None
  1366. for i, model in enumerate(model_order):
  1367. ax = fig.add_subplot(gs[0, i])
  1368. mat = np.full((n_rows, 1), np.nan)
  1369. for r, (metric_type, col) in enumerate(row_order):
  1370. sub = data[(data['Model'] == model) &
  1371. (data['metric_type'] == metric_type) &
  1372. (data['Graph'] == 'EvoAge')]
  1373. if not sub.empty:
  1374. mat[r, 0] = sub[col].values[0]
  1375. im = ax.imshow(mat, cmap=cmap, vmin=vmin, vmax=vmax, aspect='auto')
  1376. for r in range(n_rows):
  1377. if not np.isnan(mat[r, 0]):
  1378. ax.text(0, r, f'{mat[r, 0]:.2f}', ha='center', va='center', fontsize=8)
  1379. ax.axhline(3.5, color='black', linewidth=1.5)
  1380. ax.set_title(model, fontsize=11, fontweight='bold')
  1381. ax.set_xticks([0])
  1382. ax.set_xticklabels(['EvoAge'], rotation=45, ha='right', fontsize=8)
  1383. ax.set_yticks(range(n_rows))
  1384. ax.set_yticklabels(row_labels if i == 0 else [], fontsize=8)
  1385. ax.tick_params(length=0)
  1386. ax_last = fig.axes[len(model_order) - 1]
  1387. ax_last.text(0.85, 1.5, 'Validation', rotation=-90, va='center', fontsize=9)
  1388. ax_last.text(0.85, 5.5, 'Testing', rotation=-90, va='center', fontsize=9)
  1389. cax = fig.add_subplot(gs[0, -1])
  1390. cbar = fig.colorbar(im, cax=cax)
  1391. cbar.set_label('Score', fontsize=10)
  1392. plt.tight_layout()
  1393. plt.savefig('FIG2/Fig2b_heatmap.svg', dpi=300, bbox_inches='tight')
  1394. plt.show()
  1395. # %%
  1396. import pandas as pd
  1397. import numpy as np
  1398. import matplotlib.pyplot as plt
  1399. from matplotlib.colors import LinearSegmentedColormap
  1400. # ----------------------------------------------------------------------
  1401. # 1. Load data
  1402. # ----------------------------------------------------------------------
  1403. EvoAge_6model_Etype = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/training_3/Training_Stats/EvoAge_121_12M_64Emb_EdgeType_testing.csv')
  1404. EvoAge_6model_Etype
  1405. # ----------------------------------------------------------------------
  1406. # 2. Order rows by performance (RESCAL, SimplE, DistMult, RotatE, ComplEx, TransE)
  1407. # and select the columns to plot: Hit@1, Hit@3, Hit@10, MRR
  1408. # ----------------------------------------------------------------------
  1409. model_order = ['RESCAL', 'SimplE', 'DistMult', 'RotatE', 'ComplEx', 'TransE']
  1410. df = EvoAge_6model_Etype.set_index('Model Type').loc[model_order]
  1411. cols = ['Hit@1', 'Hit@3', 'Hit@10', 'MRR']
  1412. mat = df[cols].values
  1413. # ----------------------------------------------------------------------
  1414. # 3. FIGURE 2f — edge-type prediction heatmap (Hit@1, Hit@3, Hit@10, MRR)
  1415. # ----------------------------------------------------------------------
  1416. # ----- COLOR GRADIENT: 5 stops, pastel purple end -----
  1417. C1 = '#F7DFB4' # soft sand
  1418. C2 = '#F2BCBE' # soft pink
  1419. C3 = '#DBA6D2' # soft mauve
  1420. C4 = '#C79BE0' # pastel violet
  1421. C5 = '#AE8CD9' # muted purple
  1422. cmap = LinearSegmentedColormap.from_list('custom5', [
  1423. (0.00, C1),
  1424. (0.20, C2),
  1425. (0.80, C3),
  1426. (0.90, C4),
  1427. (1.00, C5),
  1428. ])
  1429. vmin, vmax = 0.0, mat.max()
  1430. fig, ax = plt.subplots(figsize=(5, 4.5))
  1431. im = ax.imshow(mat, cmap=cmap, vmin=vmin, vmax=vmax, aspect='auto')
  1432. # annotate cells
  1433. for r in range(mat.shape[0]):
  1434. for c in range(mat.shape[1]):
  1435. ax.text(c, r, f'{mat[r, c]:.2f}', ha='center', va='center', fontsize=10)
  1436. ax.set_xticks(range(len(cols)))
  1437. ax.set_xticklabels(cols, fontsize=11)
  1438. ax.set_yticks(range(len(model_order)))
  1439. ax.set_yticklabels(model_order, fontsize=11)
  1440. ax.tick_params(length=0)
  1441. # grid lines between cells
  1442. ax.set_xticks(np.arange(-0.5, len(cols), 1), minor=True)
  1443. ax.set_yticks(np.arange(-0.5, len(model_order), 1), minor=True)
  1444. ax.grid(which='minor', color='black', linewidth=1)
  1445. ax.tick_params(which='minor', length=0)
  1446. ax.set_title('Edge-type prediction', fontsize=13)
  1447. cbar = fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04)
  1448. cbar.set_label('score', fontsize=11)
  1449. plt.tight_layout()
  1450. plt.savefig('FIG2/Fig2f_edgetype-prediction.svg', dpi=300, bbox_inches='tight')
  1451. plt.show()
  1452. # %% [markdown]
  1453. # ## 2d
  1454. # %%
  1455. import pandas as pd
  1456. import numpy as np
  1457. import matplotlib.pyplot as plt
  1458. plt.rcParams['svg.fonttype'] = 'none'
  1459. # ----------------------------------------------------------------------
  1460. # 1. Load data
  1461. # ----------------------------------------------------------------------
  1462. df = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/training_3/Training_Stats/EA_A_B_Rescal_selected_metrics.csv')
  1463. # ----------------------------------------------------------------------
  1464. # 2. Filter: 121_12M ortholog, Testing split only
  1465. # ----------------------------------------------------------------------
  1466. df['Split'] = df['Split'].str.strip().str.title()
  1467. sub = df[(df['Ortholog'] == '121_12M') & (df['Split'] == 'Testing')].copy()
  1468. metrics = ['Hit@1', 'Hit@3', 'Hit@10', 'MRR']
  1469. metric_ticklabels = ['Hit@1', 'Hit@3', 'Hit@10', 'MRR']
  1470. kg_order = ['Aging', 'Biomedical', 'EvoAge']
  1471. kg_colors = {
  1472. 'Aging': '#e17f7f', # salmon/pink
  1473. 'Biomedical': '#7fb3b8', # teal
  1474. 'EvoAge': '#b67fd6', # purple
  1475. }
  1476. # ----------------------------------------------------------------------
  1477. # 3. Build plotting matrix: rows = KG, cols = metrics
  1478. # ----------------------------------------------------------------------
  1479. data = {kg: [sub[sub['KG'] == kg][m].values[0] for m in metrics] for kg in kg_order}
  1480. # ----------------------------------------------------------------------
  1481. # 4. Grouped bar plot
  1482. # ----------------------------------------------------------------------
  1483. x = np.arange(len(metrics))
  1484. n_kg = len(kg_order)
  1485. bar_width = 0.8 / n_kg
  1486. fig, ax = plt.subplots(figsize=(7, 4.5))
  1487. for i, kg in enumerate(kg_order):
  1488. offset = (i - (n_kg - 1) / 2) * bar_width
  1489. bars = ax.bar(x + offset, data[kg], width=bar_width,
  1490. color=kg_colors[kg], edgecolor='black', linewidth=0.6,
  1491. label=kg)
  1492. for rect, val in zip(bars, data[kg]):
  1493. ax.text(rect.get_x() + rect.get_width() / 2, rect.get_height() + 0.01,
  1494. f'{val:.2f}', ha='center', va='bottom', fontsize=8)
  1495. ax.set_xticks(x)
  1496. ax.set_xticklabels(metric_ticklabels, fontsize=10)
  1497. ax.set_ylabel('Score', fontsize=10)
  1498. ax.set_ylim(0, 1.05)
  1499. ax.set_title('RESCAL entity (Testing)', fontsize=11)
  1500. ax.spines['top'].set_visible(False)
  1501. ax.spines['right'].set_visible(False)
  1502. ax.legend(frameon=False, fontsize=9, loc='upper right', bbox_to_anchor=(1.15, 1.0))
  1503. plt.tight_layout()
  1504. plt.savefig('FIG2/Fig2_121_12M_testing_barplot_rescal_entity.svg', dpi=300, bbox_inches='tight')
  1505. plt.show()
  1506. # %%
  1507. import pandas as pd
  1508. import numpy as np
  1509. import matplotlib.pyplot as plt
  1510. import matplotlib as mpl
  1511. from matplotlib.patches import Rectangle
  1512. from matplotlib.lines import Line2D
  1513. mpl.rcParams['svg.fonttype'] = 'none'
  1514. # ----------------------------------------------------------------------
  1515. # 1. Load data
  1516. # ----------------------------------------------------------------------
  1517. EvoAge_Aging_biomed_edgeType = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/training_3/Training_Stats/RESCAL_all_KGtype_edgetype_results.csv')
  1518. df = EvoAge_Aging_biomed_edgeType.copy()
  1519. def get_category(name):
  1520. if name.startswith('Aging'):
  1521. return 'Aging'
  1522. elif name.startswith('Biomedical'):
  1523. return 'Biomedical'
  1524. elif name.startswith('EvoAge'):
  1525. return 'EvoAge'
  1526. def get_ortholog_type(name):
  1527. return '1:1' if '1to1' in name else '1:N'
  1528. df['Category'] = df['Dataset'].apply(get_category)
  1529. df['OrthologType'] = df['Dataset'].apply(get_ortholog_type)
  1530. cols = ['Hit@1', 'Hit@3', 'Hit@10', 'MRR']
  1531. row_order = ['1:N']
  1532. # ----------------------------------------------------------------------
  1533. # 2. Flat category colours -- sampled straight from your legend snapshot
  1534. # ----------------------------------------------------------------------
  1535. panel_colors = {
  1536. "Aging": "#e69797",
  1537. "Biomedical": "#80b4b9",
  1538. "EvoAge": "#b780d6",
  1539. }
  1540. panel_titles = {'Aging': 'Aging', 'Biomedical': 'Biomedical', 'EvoAge': 'EvoAge'}
  1541. panel_letters = {'Aging': 'b', 'Biomedical': 'c', 'EvoAge': 'd'}
  1542. # 1:1 = solid full colour, 1:N = same hue but lighter (alpha)
  1543. bar_alpha = {'1:1': 1.0, '1:N': 0.5}
  1544. row_swatch_colors = ['#8b96c9', '#e0a377'] # legend swatches (kept from original style)
  1545. # ----------------------------------------------------------------------
  1546. # 3. Build one combined figure with 3 grouped-bar panels
  1547. # ----------------------------------------------------------------------
  1548. fig, axes = plt.subplots(1, 3, figsize=(13, 3.8))
  1549. plt.subplots_adjust(wspace=0.45, top=0.72, bottom=0.28)
  1550. x = np.arange(len(cols))
  1551. bar_width = 0.5
  1552. for ax, category in zip(axes, ['Aging', 'Biomedical', 'EvoAge']):
  1553. sub = df[df['Category'] == category].set_index('OrthologType').loc[row_order]
  1554. color = panel_colors[category]
  1555. vals = sub.loc['1:N', cols].values.astype(float)
  1556. ax.bar(x, vals, width=bar_width,
  1557. color=color, alpha=bar_alpha['1:N'],
  1558. edgecolor='black', linewidth=0.8)
  1559. for xi, v in zip(x, vals):
  1560. ax.text(xi, v + 0.015, f'{v:.2f}', ha='center', va='bottom', fontsize=9)
  1561. ax.set_xticks(x)
  1562. ax.set_xticklabels(['Hit@1', 'Hit@3', 'Hit@10', 'MRR'], fontsize=10)
  1563. ax.set_ylim(0, 1.05)
  1564. ax.set_yticks([0.0, 0.2, 0.4, 0.6, 0.8, 1.0])
  1565. ax.tick_params(axis='y', labelsize=9)
  1566. ax.spines['top'].set_visible(False)
  1567. ax.spines['right'].set_visible(False)
  1568. ax.set_title(panel_titles[category], fontsize=13, pad=10)
  1569. # panel letter, top-left of panel
  1570. ax.text(-0.12, 1.08, panel_letters[category], fontsize=16, fontweight='bold',
  1571. ha='left', va='bottom', transform=ax.transAxes, clip_on=False)
  1572. # top overall title with flanking lines
  1573. fig.text(0.5, 0.92, 'RESCAL edgetype prediction', ha='center', va='center',
  1574. fontsize=14, fontweight='bold')
  1575. fig.add_artist(Line2D([0.06, 0.42], [0.93, 0.93], color='black', lw=1.3, transform=fig.transFigure))
  1576. fig.add_artist(Line2D([0.58, 0.94], [0.93, 0.93], color='black', lw=1.3, transform=fig.transFigure))
  1577. # bottom label: One-to-One plus One-to-Many (single group now, no legend split needed)
  1578. legend_y = 0.06
  1579. fig.text(0.5, legend_y, 'One-to-One plus One-to-Many', ha='center', va='center', fontsize=11)
  1580. fig.add_artist(Line2D([0.06, 0.40], [legend_y, legend_y], color='black', lw=1.3, transform=fig.transFigure))
  1581. fig.add_artist(Line2D([0.60, 0.94], [legend_y, legend_y], color='black', lw=1.3, transform=fig.transFigure))
  1582. plt.savefig('FIG2/Fig2_bcd_RESCAL_bars.svg', bbox_inches='tight')
  1583. plt.savefig('FIG2/Fig2_bcd_RESCAL_bars.png', dpi=300, bbox_inches='tight')
  1584. plt.show()
  1585. print("done")
  1586. # %% [markdown]
  1587. # ## 2e
  1588. # %%
  1589. import pandas as pd
  1590. import matplotlib.pyplot as plt
  1591. import numpy as np
  1592. # --------------------------------------------------
  1593. # Load data
  1594. # --------------------------------------------------
  1595. df = pd.read_csv("/storage/Arushi/090526_EvoAge/kg_formation/training_3/InAccuracy_format/combined_classification_accuracy_R4.csv")
  1596. graphs = [
  1597. ("Aging_121_12M", "Aging", "#e69797"),
  1598. ("Biomedical_121_12M", "Biomedical", "#80b4b9"),
  1599. ("EvoAge_121_12M", "EvoAge", "#b780d6"),
  1600. ]
  1601. metrics = ["Accuracy", "Precision", "Recall", "F1_Score", "Mean_AUC"]
  1602. labels = ["Accuracy", "Precision", "Recall", "F1", "AUC"]
  1603. plt.rcParams.update({
  1604. "font.family": "DejaVu Sans",
  1605. "font.size": 12,
  1606. "axes.linewidth": 1.2,
  1607. "svg.fonttype": "none" # keeps text editable in Illustrator/Inkscape
  1608. })
  1609. # --------------------------------------------------
  1610. # Create one figure with 3 panels
  1611. # --------------------------------------------------
  1612. fig, axes = plt.subplots(1, 3, figsize=(15, 4.5), sharey=True)
  1613. for ax, (graph, title, color) in zip(axes, graphs):
  1614. row = df[df["Graph"] == graph].iloc[0]
  1615. values = row[metrics].astype(float).values
  1616. bars = ax.bar(
  1617. np.arange(len(metrics)),
  1618. values,
  1619. width=0.62,
  1620. color=color,
  1621. edgecolor="black",
  1622. linewidth=0.8
  1623. )
  1624. # value labels
  1625. for bar, val in zip(bars, values):
  1626. ax.text(
  1627. bar.get_x() + bar.get_width()/2,
  1628. val + 0.008,
  1629. f"{val:.2f}",
  1630. ha="center",
  1631. va="bottom",
  1632. fontsize=9
  1633. )
  1634. ax.set_xticks(np.arange(len(metrics)))
  1635. ax.set_xticklabels(labels, fontsize=10)
  1636. ax.set_title(title, fontsize=15, weight="bold")
  1637. ax.grid(axis="y", linestyle="--", alpha=0.3)
  1638. ax.set_axisbelow(True)
  1639. ax.spines["top"].set_visible(False)
  1640. ax.spines["right"].set_visible(False)
  1641. axes[0].set_ylabel("Score", fontsize=12)
  1642. axes[0].set_ylim(0.60, 1.05)
  1643. # Optional overall title
  1644. # fig.suptitle("Node Classification Performance (1:1 + 1:N)", fontsize=16)
  1645. plt.tight_layout()
  1646. # --------------------------------------------------
  1647. # Save ONE SVG and ONE PNG
  1648. # --------------------------------------------------
  1649. plt.savefig(
  1650. "FIG2/triple_classification_KG_Accuracy_1to1plus1toMany.svg",
  1651. format="svg",
  1652. bbox_inches="tight"
  1653. )
  1654. plt.savefig(
  1655. "FIG2/_TRIPLE_CLASSIFICATION_KG_Accuracy_1to1plus1toMany.png",
  1656. dpi=600,
  1657. bbox_inches="tight"
  1658. )
  1659. plt.show()
  1660. # %% [markdown]
  1661. # ## 2f
  1662. # %%
  1663. # Edge_type = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/training__2/Shuffled_EvoAge_testing/Store_House/shuffled_test_sets/rescal_shuffled_metrics.csv')
  1664. # Edge_type
  1665. #!/usr/bin/env python3
  1666. """
  1667. Monte Carlo P-Value Analysis — Radial (Spider) Plot
  1668. Compare real EvoAge test set performance vs 30 shuffled test sets.
  1669. Produces a single radial/spider chart (matching the reference figure):
  1670. - Red polygon = real EvoAge test performance
  1671. - Blue polygon = mean of 30 shuffled null-hypothesis baselines
  1672. - Value labels on every spoke
  1673. - "Statistical Significance" box (per-metric Monte Carlo p-value + marker)
  1674. - "Interpretation" box
  1675. Output is saved as a PDF (vector, publication-ready).
  1676. """
  1677. import pandas as pd
  1678. import numpy as np
  1679. import matplotlib.pyplot as plt
  1680. from scipy import stats
  1681. plt.rcParams.update({
  1682. "font.family": "DejaVu Sans",
  1683. "svg.fonttype": "none"
  1684. })
  1685. # ---- Your real test set results ----
  1686. REAL_RESULTS = {
  1687. 'HITS@1': 0.810538,
  1688. 'HITS@3': 0.924235,
  1689. 'HITS@10': 0.997215,
  1690. 'MR': 1.520212,
  1691. 'MRR': 0.875109
  1692. }
  1693. # ---- Load your shuffled results ----
  1694. results_df = pd.read_csv(
  1695. '/storage/Arushi/090526_EvoAge/kg_formation/training_3/Shuffled_EvoAge_testing/Store_House/shuffled_test_sets/rescal_shuffled_metrics.csv'
  1696. )
  1697. print("=" * 70)
  1698. print("MONTE CARLO P-VALUE ANALYSIS")
  1699. print("Real Test Set vs 30 Shuffled Test Sets")
  1700. print("=" * 70)
  1701. print(f"\n[1/3] Loaded {len(results_df)} shuffled test set results")
  1702. print(results_df[['shuffle_id', 'MRR', 'HITS@1', 'HITS@10']].head(10))
  1703. # Order controls the layout of the radial chart (clockwise from the top)
  1704. metrics = ['MRR', 'HITS@1', 'HITS@3', 'HITS@10']
  1705. # For MRR, HITS@K: higher = better -> test if real is significantly HIGHER
  1706. # For MR: lower = better -> test if real is significantly LOWER
  1707. higher_is_better = {'MRR': True, 'HITS@1': True, 'HITS@3': True, 'HITS@10': True}
  1708. # ---- Compute p-values ----
  1709. print("\n[2/3] Computing Monte Carlo p-values...")
  1710. summary_rows = []
  1711. for metric in metrics:
  1712. shuffled_vals = results_df[metric].dropna().values
  1713. real_val = REAL_RESULTS[metric]
  1714. n = len(shuffled_vals)
  1715. shuffled_mean = shuffled_vals.mean()
  1716. shuffled_std = shuffled_vals.std(ddof=1)
  1717. shuffled_min = shuffled_vals.min()
  1718. shuffled_max = shuffled_vals.max()
  1719. # ---- MONTE CARLO P-VALUE (empirical) ----
  1720. if higher_is_better[metric]:
  1721. n_extreme = np.sum(shuffled_vals >= real_val)
  1722. else:
  1723. n_extreme = np.sum(shuffled_vals <= real_val)
  1724. # Empirical p-value with +1/+1 correction (standard in permutation testing)
  1725. monte_carlo_p = (n_extreme + 1) / (n + 1)
  1726. # ---- Z-TEST P-VALUE (parametric, for reference) ----
  1727. z_score = (real_val - shuffled_mean) / shuffled_std if shuffled_std > 0 else np.nan
  1728. if higher_is_better[metric]:
  1729. z_p = stats.norm.sf(z_score)
  1730. else:
  1731. z_p = stats.norm.cdf(z_score)
  1732. # Cohen's d (effect size)
  1733. cohens_d = (real_val - shuffled_mean) / shuffled_std if shuffled_std > 0 else np.nan
  1734. summary_rows.append({
  1735. 'metric': metric,
  1736. 'real_value': real_val,
  1737. 'shuffled_mean': shuffled_mean,
  1738. 'shuffled_std': shuffled_std,
  1739. 'shuffled_min': shuffled_min,
  1740. 'shuffled_max': shuffled_max,
  1741. 'n_shuffles': n,
  1742. 'n_extreme': n_extreme,
  1743. 'z_score': z_score,
  1744. 'cohens_d': cohens_d,
  1745. 'p_monte_carlo': monte_carlo_p,
  1746. 'p_ztest': z_p
  1747. })
  1748. summary_df = pd.DataFrame(summary_rows)
  1749. pd.set_option('display.float_format', lambda x: f'{x:.6g}')
  1750. pd.set_option('display.max_columns', None)
  1751. pd.set_option('display.width', None)
  1752. print("\n" + summary_df.to_string(index=False))
  1753. # ---- Save results table ----
  1754. output_dir = 'FIG2/'
  1755. summary_df.to_csv(f'{output_dir}/monte_carlo_pvalues_vs_shuffled.csv', index=False)
  1756. print(f"\n✓ Saved to: {output_dir}/monte_carlo_pvalues_vs_shuffled.csv")
  1757. # ---- Radial (spider) plot ----
  1758. print("\n[3/3] Creating radial visualization...")
  1759. def sig_marker(p):
  1760. if p < 0.001:
  1761. return '***'
  1762. if p < 0.01:
  1763. return '**'
  1764. if p < 0.05:
  1765. return '*'
  1766. return 'n.s.'
  1767. real_vals = [summary_df.loc[summary_df['metric'] == m, 'real_value'].values[0] for m in metrics]
  1768. shuf_vals = [summary_df.loc[summary_df['metric'] == m, 'shuffled_mean'].values[0] for m in metrics]
  1769. N = len(metrics)
  1770. angles = np.linspace(0, 2 * np.pi, N, endpoint=False).tolist()
  1771. real_plot = real_vals + real_vals[:1]
  1772. shuf_plot = shuf_vals + shuf_vals[:1]
  1773. angles_plot = angles + angles[:1]
  1774. fig = plt.figure(figsize=(11, 10))
  1775. ax = fig.add_subplot(111, polar=True)
  1776. ax.set_theta_zero_location('N') # first metric (MRR) at the top
  1777. ax.set_theta_direction(-1) # go clockwise
  1778. ax.set_xticks(angles)
  1779. ax.set_xticklabels(metrics, fontsize=12, fontweight='bold')
  1780. max_r = max(real_vals + shuf_vals) * 1.25
  1781. ax.set_ylim(0, max_r)
  1782. ax.set_rgrids(np.linspace(max_r / 5, max_r, 5), angle=0, fontsize=8)
  1783. ax.grid(alpha=0.4, linestyle='--')
  1784. RED = '#d62728'
  1785. BLUE = '#4a78c9'
  1786. # Real EvoAge performance
  1787. ax.plot(angles_plot, real_plot, color=RED, linewidth=2.2, label='EvoAge test')
  1788. ax.fill(angles_plot, real_plot, color=RED, alpha=0.15)
  1789. ax.scatter(angles, real_vals, color=RED, s=45, zorder=5)
  1790. # Shuffled null baseline (mean of 30 shuffles)
  1791. ax.plot(angles_plot, shuf_plot, color=BLUE, linewidth=2.2, label='EvoAge Shuffled test')
  1792. ax.fill(angles_plot, shuf_plot, color=BLUE, alpha=0.12)
  1793. ax.scatter(angles, shuf_vals, color=BLUE, s=45, zorder=5)
  1794. # Value labels
  1795. for ang, val in zip(angles, real_vals):
  1796. ax.annotate(f'{val:.4f}', xy=(ang, val), xytext=(ang, val + max_r * 0.03),
  1797. ha='center', va='bottom', fontsize=9, fontweight='bold', color=RED,
  1798. bbox=dict(boxstyle='round,pad=0.25', fc='white', ec=RED, lw=1))
  1799. for ang, val in zip(angles, shuf_vals):
  1800. ax.annotate(f'{val:.4f}', xy=(ang, val), xytext=(ang, val - max_r * 0.03),
  1801. ha='center', va='top', fontsize=8, style='italic', color=BLUE,
  1802. bbox=dict(boxstyle='round,pad=0.2', fc='white', ec=BLUE, lw=0.8))
  1803. ax.set_title(
  1804. 'RESCAL KG Embedding: Monte Carlo Null Hypothesis Testing\n'
  1805. 'Real Performance vs 30 Shuffled Test Set Baselines (Radial View)',
  1806. fontsize=13, fontweight='bold', pad=30
  1807. )
  1808. ax.legend(loc='lower right', bbox_to_anchor=(1.3, -0.05), fontsize=12, frameon=False)
  1809. # Statistical significance box
  1810. sig_lines = "\n".join(
  1811. f"{row['metric']}: {sig_marker(row['p_monte_carlo'])} (p={row['p_monte_carlo']:.3g})"
  1812. for _, row in summary_df.iterrows()
  1813. )
  1814. fig.text(0.03, 0.35,
  1815. f"Statistical Significance:\n\n{sig_lines}",
  1816. fontsize=9, family='monospace',
  1817. bbox=dict(boxstyle='round,pad=0.6', fc='#fdf6d8', ec='#c9b95a', lw=1))
  1818. # Interpretation box
  1819. interp_text = (
  1820. "Interpretation:\n"
  1821. "- Red polygon = Real EvoAge KG performance\n"
  1822. "- Blue polygon = Shuffled null baseline (mean)\n"
  1823. "- Larger red area beyond blue = significant gain\n"
  1824. "- Values labeled on each spoke"
  1825. )
  1826. fig.text(0.03, 0.12, interp_text, fontsize=9,
  1827. bbox=dict(boxstyle='round,pad=0.6', fc='#e4f2fb', ec='#7fb3d9', lw=1))
  1828. fig.tight_layout()
  1829. pdf_path = f'{output_dir}/monte_carlo_radial_comparison.pdf'
  1830. fig.savefig(pdf_path, format='pdf', bbox_inches='tight')
  1831. print(f"✓ Saved radial plot to: {pdf_path}")
  1832. plt.show()
  1833. plt.close(fig)
  1834. # ---- Final interpretation (console) ----
  1835. print("\n" + "=" * 70)
  1836. print("INTERPRETATION")
  1837. print("=" * 70)
  1838. for _, row in summary_df.iterrows():
  1839. metric = row['metric']
  1840. real = row['real_value']
  1841. shuffled_mean = row['shuffled_mean']
  1842. p_val = row['p_monte_carlo']
  1843. cohens_d = row['cohens_d']
  1844. if p_val < 0.001:
  1845. sig_text = "*** HIGHLY SIGNIFICANT"
  1846. elif p_val < 0.01:
  1847. sig_text = "** VERY SIGNIFICANT"
  1848. elif p_val < 0.05:
  1849. sig_text = "* SIGNIFICANT"
  1850. else:
  1851. sig_text = "n.s. NOT SIGNIFICANT"
  1852. ratio = real / shuffled_mean if shuffled_mean != 0 else np.inf
  1853. print(f"\n{metric:10s}")
  1854. print(f" Real: {real:.6f}")
  1855. print(f" Shuffled mean: {shuffled_mean:.6f} (ratio: {ratio:.2f}x)")
  1856. print(f" p-value: {p_val:.6f} {sig_text}")
  1857. print(f" Effect size: Cohen's d = {cohens_d:+.3f}")
  1858. print("\n" + "=" * 70)
  1859. # %% [markdown]
  1860. # ## 2h
  1861. # %%
  1862. import pandas as pd
  1863. import matplotlib.pyplot as plt
  1864. import numpy as np
  1865. plt.rcParams.update({
  1866. "font.family": "DejaVu Sans",
  1867. "svg.fonttype": "none"
  1868. })
  1869. k = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/building_evoage_with_1_percent_species_testsplit/1_per_test_triple_count.csv')
  1870. print(k.columns.tolist())
  1871. # Color scheme matching original figure
  1872. colors = {
  1873. 'Human': '#A98BC9',
  1874. 'Mouse': '#8AD3E4',
  1875. 'Celegans': '#C9C97E',
  1876. 'Drosophila': '#8FCB7E',
  1877. 'Zebrafish': '#E89BC0',
  1878. 'Yeast': '#C9A98B',
  1879. 'CrossSpecies': '#F7DADF',
  1880. }
  1881. order = ['Human', 'Mouse', 'Celegans', 'Drosophila', 'Zebrafish', 'Yeast', 'CrossSpecies']
  1882. # ---------- Panel K: donut chart of heldout data (log-scaled wedges, real labels) ----------
  1883. k = k.set_index('Species').loc[order].reset_index()
  1884. # Log-transform wedge sizes so small species are still visible,
  1885. # but keep the real counts as text labels.
  1886. log_vals = np.log10(k['Line_Count'])
  1887. log_vals = log_vals - log_vals.min() + 1 # shift so smallest still gets a visible wedge
  1888. fig, ax = plt.subplots(figsize=(6, 6))
  1889. wedges, _ = ax.pie(
  1890. log_vals,
  1891. colors=[colors[s] for s in k['Species']],
  1892. startangle=90,
  1893. counterclock=False,
  1894. wedgeprops=dict(width=0.35, edgecolor='white', linewidth=1.5)
  1895. )
  1896. total = k['Line_Count'].sum()
  1897. ax.text(0, 0.05, 'total', ha='center', va='center', fontsize=13, fontweight='bold')
  1898. ax.text(0, -0.08, f'{total:,}', ha='center', va='center', fontsize=12)
  1899. # Labels read back from each wedge's true midpoint angle -> can never drift onto the wrong wedge
  1900. for w, (_, row) in zip(wedges, k.iterrows()):
  1901. ang = np.radians((w.theta1 + w.theta2) / 2)
  1902. x = 1.3 * np.cos(ang)
  1903. y = 1.3 * np.sin(ang)
  1904. ax.text(x, y, f"{row['Species']}\n{row['Line_Count']:,}",
  1905. ha='center', va='center', fontsize=8)
  1906. ax.set_title('heldout data (wedges log-scaled; labels = actual counts)',
  1907. fontsize=12, fontweight='bold', loc='left')
  1908. plt.tight_layout()
  1909. plt.savefig('FIG2/fig2n.pdf')
  1910. plt.show()
  1911. plt.close()
  1912. k.head()
  1913. # %% [markdown]
  1914. # ## 2i
  1915. # %%
  1916. import pandas as pd
  1917. import matplotlib.pyplot as plt
  1918. import numpy as np
  1919. plt.rcParams.update({
  1920. "font.family": "DejaVu Sans",
  1921. "svg.fonttype": "none"
  1922. })
  1923. # ================================================================
  1924. # CANONICAL PALETTE (matches icon legend)
  1925. # ================================================================
  1926. colors = {
  1927. 'Human': '#A98BC9',
  1928. 'Mouse': '#8AD3E4',
  1929. 'Celegans': '#C9C97E',
  1930. 'Drosophila': '#8FCB7E',
  1931. 'Zebrafish': '#E89BC0',
  1932. 'Yeast': '#C9A98B',
  1933. 'CrossSpecies': '#F7DADF',
  1934. }
  1935. order = ['Human', 'Mouse', 'Celegans', 'Drosophila', 'Zebrafish', 'Yeast', 'CrossSpecies']
  1936. # ================================================================
  1937. # FIG 2N — heldout-data donut
  1938. # ================================================================
  1939. def make_donut():
  1940. k = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/'
  1941. 'building_evoage_with_1_percent_species_testsplit/1_per_test_triple_count.csv')
  1942. k = k.set_index('Species').loc[order].reset_index()
  1943. log_vals = np.log10(k['Line_Count'])
  1944. log_vals = log_vals - log_vals.min() + 1 # smallest wedge stays visible
  1945. fig, ax = plt.subplots(figsize=(6, 6))
  1946. wedges, _ = ax.pie(
  1947. log_vals,
  1948. colors=[colors[s] for s in k['Species']],
  1949. startangle=90,
  1950. counterclock=False,
  1951. wedgeprops=dict(width=0.35, edgecolor='white', linewidth=1.5)
  1952. )
  1953. total = k['Line_Count'].sum()
  1954. ax.text(0, 0.05, 'total', ha='center', va='center', fontsize=13, fontweight='bold')
  1955. ax.text(0, -0.08, f'{total:,}', ha='center', va='center', fontsize=12)
  1956. # labels read from each wedge's true midpoint -> never drift onto wrong wedge
  1957. for w, (_, row) in zip(wedges, k.iterrows()):
  1958. ang = np.radians((w.theta1 + w.theta2) / 2)
  1959. x, y = 1.3 * np.cos(ang), 1.3 * np.sin(ang)
  1960. ax.text(x, y, f"{row['Species']}\n{row['Line_Count']:,}",
  1961. ha='center', va='center', fontsize=8)
  1962. ax.set_title('heldout data (wedges log-scaled; labels = actual counts)',
  1963. fontsize=12, fontweight='bold', loc='left')
  1964. plt.tight_layout()
  1965. plt.savefig('FIG2/fig2n.pdf', bbox_inches='tight')
  1966. plt.close()
  1967. # ================================================================
  1968. # FIG 2O — link-prediction hit@k grouped bars
  1969. # ================================================================
  1970. def make_linkpred():
  1971. l = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/training_3/'
  1972. '1_per_species_test_EvoAge_121_12M/rescal_1_percent_test_linkprediction_results.csv')
  1973. l['TestSet'] = l['TestSet'].replace({'Unknown_0': 'Human', 'Unknown_6': 'CrossSpecies'})
  1974. l = l.set_index('TestSet').loc[order].reset_index() # ordered; index i == order[i]
  1975. hits = ['Hits@1', 'Hits@3', 'Hits@10']
  1976. x_labels = ['1', '3', '10']
  1977. n_species = len(order)
  1978. width = 0.8 / n_species
  1979. x = np.arange(len(hits))
  1980. fig, ax = plt.subplots(figsize=(6, 5))
  1981. for i, sp in enumerate(order):
  1982. vals = l.iloc[i][hits].values.astype(float) # by position -> no filter mismatch
  1983. ax.bar(x + (i - n_species / 2) * width + width / 2, vals,
  1984. width=width, color=colors[sp], label=sp)
  1985. ax.set_ylim(0, 1.05)
  1986. ax.set_xticks(x)
  1987. ax.set_xticklabels(x_labels)
  1988. ax.set_xlabel('hit@')
  1989. ax.set_ylabel('score')
  1990. ax.set_title('prediction (link)', fontsize=14, fontweight='bold')
  1991. ax.legend(bbox_to_anchor=(1.01, 1), loc='upper left', fontsize=8, frameon=False)
  1992. ax.spines['top'].set_visible(False)
  1993. ax.spines['right'].set_visible(False)
  1994. plt.tight_layout()
  1995. plt.savefig('FIG2/fig2o_new.pdf', bbox_inches='tight')
  1996. plt.show()
  1997. plt.close()
  1998. if __name__ == '__main__':
  1999. make_donut()
  2000. make_linkpred()
  2001. print('done')
  2002. # %%
  2003. import pandas as pd
  2004. import matplotlib.pyplot as plt
  2005. import numpy as np
  2006. plt.rcParams.update({
  2007. "font.family": "DejaVu Sans",
  2008. "svg.fonttype": "none"
  2009. })
  2010. # ================================================================
  2011. # CANONICAL PALETTE (matches icon legend)
  2012. # ================================================================
  2013. colors = {
  2014. 'Human': '#A98BC9',
  2015. 'Mouse': '#8AD3E4',
  2016. 'Celegans': '#C9C97E',
  2017. 'Drosophila': '#8FCB7E',
  2018. 'Zebrafish': '#E89BC0',
  2019. 'Yeast': '#C9A98B',
  2020. 'CrossSpecies': '#F7DADF',
  2021. }
  2022. order = ['Human', 'Mouse', 'Celegans', 'Drosophila', 'Zebrafish', 'Yeast', 'CrossSpecies']
  2023. # ================================================================
  2024. # FIG 2N — heldout-data donut
  2025. # ================================================================
  2026. def make_donut():
  2027. k = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/'
  2028. 'building_evoage_with_1_percent_species_testsplit/1_per_test_triple_count.csv')
  2029. k = k.set_index('Species').loc[order].reset_index()
  2030. log_vals = np.log10(k['Line_Count'])
  2031. log_vals = log_vals - log_vals.min() + 1 # smallest wedge stays visible
  2032. fig, ax = plt.subplots(figsize=(6, 6))
  2033. wedges, _ = ax.pie(
  2034. log_vals,
  2035. colors=[colors[s] for s in k['Species']],
  2036. startangle=90,
  2037. counterclock=False,
  2038. wedgeprops=dict(width=0.35, edgecolor='white', linewidth=1.5)
  2039. )
  2040. total = k['Line_Count'].sum()
  2041. ax.text(0, 0.05, 'total', ha='center', va='center', fontsize=13, fontweight='bold')
  2042. ax.text(0, -0.08, f'{total:,}', ha='center', va='center', fontsize=12)
  2043. # labels read from each wedge's true midpoint -> never drift onto wrong wedge
  2044. for w, (_, row) in zip(wedges, k.iterrows()):
  2045. ang = np.radians((w.theta1 + w.theta2) / 2)
  2046. x, y = 1.3 * np.cos(ang), 1.3 * np.sin(ang)
  2047. ax.text(x, y, f"{row['Species']}\n{row['Line_Count']:,}",
  2048. ha='center', va='center', fontsize=8)
  2049. ax.set_title('heldout data (wedges log-scaled; labels = actual counts)',
  2050. fontsize=12, fontweight='bold', loc='left')
  2051. plt.tight_layout()
  2052. plt.savefig('FIG2/fig2n.pdf', bbox_inches='tight')
  2053. plt.close()
  2054. # ================================================================
  2055. # FIG 2O — link-prediction hit@k grouped bars
  2056. # ================================================================
  2057. def make_linkpred():
  2058. l = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/training_3/'
  2059. '1_per_species_test_EvoAge_121_12M/rescal_1_percent_test_linkprediction_results.csv')
  2060. l['TestSet'] = l['TestSet'].replace({'Unknown_0': 'Human', 'Unknown_6': 'CrossSpecies'})
  2061. l = l.set_index('TestSet').loc[order].reset_index() # ordered; index i == order[i]
  2062. hits = ['Hits@1', 'Hits@3', 'Hits@10']
  2063. x_labels = ['1', '3', '10']
  2064. n_species = len(order)
  2065. width = 0.8 / n_species
  2066. x = np.arange(len(hits))
  2067. fig, ax = plt.subplots(figsize=(6, 5))
  2068. for i, sp in enumerate(order):
  2069. vals = l.iloc[i][hits].values.astype(float) # by position -> no filter mismatch
  2070. ax.bar(x + (i - n_species / 2) * width + width / 2, vals,
  2071. width=width, color=colors[sp], label=sp)
  2072. ax.set_ylim(0, 1.05)
  2073. ax.set_xticks(x)
  2074. ax.set_xticklabels(x_labels)
  2075. ax.set_xlabel('hit@')
  2076. ax.set_ylabel('score')
  2077. ax.set_title('prediction (link)', fontsize=14, fontweight='bold')
  2078. ax.legend(bbox_to_anchor=(1.01, 1), loc='upper left', fontsize=8, frameon=False)
  2079. ax.spines['top'].set_visible(False)
  2080. ax.spines['right'].set_visible(False)
  2081. plt.tight_layout()
  2082. plt.savefig('FIG2/fig2o_new.pdf', bbox_inches='tight')
  2083. plt.close()
  2084. # ================================================================
  2085. # FIG 2P — edge-type prediction hit@k line plot
  2086. # ================================================================
  2087. def make_edgetype():
  2088. m = pd.read_csv('/storage/Arushi/090526_EvoAge/kg_formation/training_3/'
  2089. '1_per_species_test_EvoAge_121_12M/rescal_1_percent_test_edgeType_results.csv')
  2090. m['TestSet'] = m['TestSet'].replace({'Unknown_0': 'Human', 'Unknown_6': 'CrossSpecies'})
  2091. print(m['TestSet'].unique())
  2092. hits = ['Hits@1', 'Hits@3', 'Hits@10']
  2093. available = m['TestSet'].unique()
  2094. order_m = [s for s in order if s in available] # keep canonical order, drop missing
  2095. fig, ax = plt.subplots(figsize=(6, 5))
  2096. x_pos = [0, 1, 2] # evenly spaced positions so '3' sits centered, not skewed toward '1'
  2097. for sp in order_m:
  2098. vals = m.loc[m['TestSet'] == sp, hits].values.flatten().astype(float)
  2099. ax.plot(x_pos, vals, marker='o', color=colors[sp], label=sp, linewidth=2)
  2100. ax.set_ylim(0, 1.05)
  2101. ax.set_xticks(x_pos)
  2102. ax.set_xticklabels(['1', '3', '10'])
  2103. ax.set_xlabel('hit@')
  2104. ax.set_ylabel('score')
  2105. ax.set_title('prediction (edge type)', fontsize=14, fontweight='bold')
  2106. ax.legend(bbox_to_anchor=(1.01, 1), loc='upper left', fontsize=8, frameon=False)
  2107. ax.spines['top'].set_visible(False)
  2108. ax.spines['right'].set_visible(False)
  2109. plt.tight_layout()
  2110. plt.savefig('FIG2/fig2p_new.pdf', bbox_inches='tight')
  2111. plt.show()
  2112. plt.close()
  2113. if __name__ == '__main__':
  2114. make_donut()
  2115. make_linkpred()
  2116. make_edgetype()
  2117. print('done')
  2118. # %% [markdown]
  2119. # ## 2j
  2120. # %%
  2121. import pandas as pd
  2122. import numpy as np
  2123. import matplotlib.pyplot as plt
  2124. plt.rcParams.update({
  2125. "font.family": "DejaVu Sans",
  2126. "svg.fonttype": "none"
  2127. })
  2128. # ----------------------------------------------------------------------
  2129. # 1. Load data
  2130. # ----------------------------------------------------------------------
  2131. dataframe = pd.read_csv(
  2132. "/storage/Arushi/090526_EvoAge/kg_formation/training_3/Agingonly_testing_data/Aging_Only_Testingdata_Evaluated_onall_rescal_trainedKG's_summary.csv"
  2133. )
  2134. # ----------------------------------------------------------------------
  2135. # 2. Keep only 1:N (121_12M)
  2136. # ----------------------------------------------------------------------
  2137. dataframe = dataframe[dataframe["Ortholog"] == "121_12M"].copy()
  2138. ortholog_label = {"121_12M": "1:N"}
  2139. dataframe["OrthologLabel"] = dataframe["Ortholog"].map(ortholog_label)
  2140. dataframe["Condition"] = dataframe["Graph"] + " (" + dataframe["OrthologLabel"] + ")"
  2141. # Order
  2142. condition_order = [
  2143. "Aging -> Aging (1:N)",
  2144. "Aging -> EvoAge (1:N)",
  2145. ]
  2146. dataframe["Condition"] = pd.Categorical(
  2147. dataframe["Condition"],
  2148. categories=condition_order,
  2149. ordered=True,
  2150. )
  2151. dataframe = dataframe.sort_values("Condition")
  2152. # Colours
  2153. condition_colors = {
  2154. "Aging -> Aging (1:N)": "#e07a7a",
  2155. "Aging -> EvoAge (1:N)": "#4f8fd9",
  2156. }
  2157. metrics = ["Hit@1", "Hit@3", "Hit@10", "MRR"]
  2158. # ----------------------------------------------------------------------
  2159. # 3. Grouped bar chart
  2160. # ----------------------------------------------------------------------
  2161. fig, ax = plt.subplots(figsize=(6, 5))
  2162. n_conditions = len(condition_order)
  2163. width = 0.35
  2164. x = np.arange(len(metrics))
  2165. for i, cond in enumerate(condition_order):
  2166. row = dataframe[dataframe["Condition"] == cond]
  2167. if row.empty:
  2168. continue
  2169. vals = row[metrics].values.flatten()
  2170. offset = (i - (n_conditions - 1) / 2) * width
  2171. ax.bar(
  2172. x + offset,
  2173. vals,
  2174. width,
  2175. label=cond,
  2176. color=condition_colors[cond],
  2177. )
  2178. ax.set_xticks(x)
  2179. ax.set_xticklabels(metrics, fontsize=11)
  2180. ax.set_ylabel("Score", fontsize=12)
  2181. ax.set_ylim(0, 1.05)
  2182. ax.legend(frameon=False, fontsize=10)
  2183. ax.spines["top"].set_visible(False)
  2184. ax.spines["right"].set_visible(False)
  2185. plt.tight_layout()
  2186. plt.savefig("FIG2/2oAging_crossKG_barplot_1N_only.pdf", dpi=300, bbox_inches="tight")
  2187. plt.show()
  2188. # %% [markdown]
  2189. # ## Main Figure 4
  2190. # %% [markdown]
  2191. # ## 4b
  2192. # %%
  2193. from __future__ import annotations
  2194. %matplotlib inline
  2195. import csv
  2196. import re
  2197. import shutil
  2198. import urllib.parse
  2199. import zipfile
  2200. from collections import Counter, defaultdict
  2201. from pathlib import Path
  2202. import numpy as np
  2203. import matplotlib.pyplot as plt
  2204. from matplotlib.colors import LinearSegmentedColormap
  2205. from matplotlib.backends.backend_pdf import PdfPages
  2206. # ============================================================
  2207. # Configuration
  2208. # ============================================================
  2209. INPUT = Path('/storage/Arushi/090526_EvoAge/multiagent_hypo/evoage_100_right_inverse/all_LLMs_hypothesis_response.csv')
  2210. ROOT = Path('/storage/Arushi/090526_EvoAge/multiagent_hypo/evoage_100_right_inverse/My_final_analysis/My_analysis_doi_paired_retention_plots')
  2211. MASTER_PDF = ROOT / 'retention_pie_donut_plots.pdf'
  2212. ALL_MODELS = ['EvoAge', 'medgemma', 'BioMistral', 'BioMistralFinetuned']
  2213. DISPLAY = {
  2214. 'EvoAge': 'EvoAge',
  2215. 'medgemma': 'MedGemma',
  2216. 'BioMistral': 'BioMistral',
  2217. 'BioMistralFinetuned': 'BioMistral fine-tuned',
  2218. }
  2219. PRIMARY_MODELS = ['EvoAge', 'medgemma', 'BioMistral']
  2220. ALT_MODELS = ['EvoAge', 'BioMistral', 'BioMistralFinetuned']
  2221. TYPES = ['Right', 'Inverse']
  2222. VERDICTS = ['no_support', 'weak_support', 'partial_support', 'support', 'strong_support']
  2223. VERDICT_LABELS = ['No support', 'Weak support', 'Partial support', 'Support', 'Strong support']
  2224. SCORE = {'no_support': 0, 'weak_support': 1, 'partial_support': 2, 'support': 3, 'strong_support': 4}
  2225. UNKNOWN = {'unknown', 'missing', 'unknown/missing', 'unknown / missing', ''}
  2226. MODEL_COLORS = {'EvoAge': '#4C72B0', 'medgemma': '#DD8452', 'BioMistral': '#55A868',
  2227. 'BioMistralFinetuned': '#937860'}
  2228. KEPT_COLOR = '#4C9F70'
  2229. REMOVED_COLOR = '#C44E52'
  2230. MULTI_COLOR = '#8172B3'
  2231. BAR_COLOR = '#1f77b4'
  2232. GRADIENT_ENDPOINTS = ['#1a9850', '#d61f96']
  2233. _ordinal_cmap = LinearSegmentedColormap.from_list('green_to_magenta', GRADIENT_ENDPOINTS)
  2234. VERDICT_COLORS = [_ordinal_cmap(i / (len(VERDICTS) - 1)) for i in range(len(VERDICTS))]
  2235. plt.rcParams.update({
  2236. 'font.family': 'DejaVu Sans',
  2237. 'font.size': 9,
  2238. 'axes.titlesize': 11,
  2239. 'axes.labelsize': 9,
  2240. 'legend.fontsize': 8,
  2241. 'axes.linewidth': 0.8,
  2242. 'patch.linewidth': 0.7,
  2243. 'axes.spines.top': False,
  2244. 'axes.spines.right': False,
  2245. })
  2246. # ============================================================
  2247. # Helper Functions
  2248. # ============================================================
  2249. def normalize_doi(value: str) -> str:
  2250. value = (value or '').strip().lower()
  2251. value = re.sub(r'^https?://(dx\.)?doi\.org/', '', value)
  2252. value = re.sub(r'^doi:\s*', '', value)
  2253. value = urllib.parse.unquote(value).strip().rstrip(' .;,')
  2254. return value
  2255. def normalize_verdict(value: str) -> str:
  2256. value = (value or '').strip().lower().replace('-', '_').replace(' ', '_')
  2257. value = re.sub(r'_+', '_', value)
  2258. aliases = {'nosupport': 'no_support', 'partialsupport': 'partial_support',
  2259. 'weaksupport': 'weak_support', 'strongsupport': 'strong_support',
  2260. 'unknown_missing': 'unknown'}
  2261. return aliases.get(value, value)
  2262. def footer(fig, text):
  2263. fig.text(0.01, 0.01, text, ha='left', va='bottom', fontsize=6.5)
  2264. def pct_count_autopct(total: int):
  2265. def _fmt(pct):
  2266. count = int(round(pct * total / 100.0))
  2267. return f'{pct:.1f}%\n(n={count})'
  2268. return _fmt
  2269. # ============================================================
  2270. # Load and Process Data
  2271. # ============================================================
  2272. with INPUT.open(newline='', encoding='utf-8-sig') as handle:
  2273. reader = csv.DictReader(handle)
  2274. raw = []
  2275. for row in reader:
  2276. model = row['model'].strip()
  2277. if model not in ALL_MODELS:
  2278. continue
  2279. raw.append({
  2280. 'doi': normalize_doi(row['DOI']),
  2281. 'model': model,
  2282. 'hypothesis_type': row['hypothesis_type'].strip().title(),
  2283. 'verdict': normalize_verdict(row['verdict']),
  2284. })
  2285. by_doi = defaultdict(list)
  2286. for row in raw:
  2287. by_doi[row['doi']].append(row)
  2288. dois = sorted(by_doi)
  2289. lookup_all = {(r['doi'], r['model'], r['hypothesis_type']): r for r in raw}
  2290. def scores(model: str, direction: str, doi_list) -> np.ndarray:
  2291. return np.array([SCORE[lookup_all[(doi, model, direction)]['verdict']] for doi in doi_list], dtype=float)
  2292. # Filtering
  2293. trigger_rows = [r for r in raw if r['model'] in set(ALL_MODELS) and r['verdict'] in UNKNOWN]
  2294. excluded_dois = sorted({r['doi'] for r in trigger_rows})
  2295. retained_dois = [doi for doi in dois if doi not in set(excluded_dois)]
  2296. n_total = len(dois)
  2297. n_removed = len(excluded_dois)
  2298. n_kept = len(retained_dois)
  2299. flagged_by = defaultdict(set)
  2300. for r in trigger_rows:
  2301. flagged_by[r['doi']].add(r['model'])
  2302. per_model_removed = {m: len({r['doi'] for r in trigger_rows if r['model'] == m}) for m in PRIMARY_MODELS}
  2303. per_model_kept = {m: n_total - per_model_removed[m] for m in PRIMARY_MODELS}
  2304. attribution = Counter()
  2305. for doi, models in flagged_by.items():
  2306. if len(models) == 1:
  2307. attribution[next(iter(models))] += 1
  2308. else:
  2309. attribution['Multiple'] += 1
  2310. print(f"Data loaded: {n_total} total DOIs, {n_kept} retained, {n_removed} removed")
  2311. # %%
  2312. # ============================================================
  2313. # Plot 3: Per-model % kept vs % removed (Donut Charts)
  2314. # ============================================================
  2315. fig, axes = plt.subplots(1, len(PRIMARY_MODELS), figsize=(4.4 * len(PRIMARY_MODELS), 4.6))
  2316. for ax, model in zip(axes, PRIMARY_MODELS):
  2317. ax.pie([per_model_kept[model], per_model_removed[model]],
  2318. colors=[KEPT_COLOR, REMOVED_COLOR],
  2319. autopct=pct_count_autopct(n_total), startangle=90, counterclock=False,
  2320. wedgeprops=dict(width=0.42, edgecolor='white'), pctdistance=0.75,
  2321. textprops=dict(fontsize=8))
  2322. ax.text(0, 0, DISPLAY[model], ha='center', va='center', fontsize=10, fontweight='bold')
  2323. ax.set_title(f'{DISPLAY[model]}\nkept vs removed', fontsize=10)
  2324. fig.subplots_adjust(bottom=0.2)
  2325. handles = [plt.Rectangle((0, 0), 1, 1, color=KEPT_COLOR),
  2326. plt.Rectangle((0, 0), 1, 1, color=REMOVED_COLOR)]
  2327. fig.legend(handles, ['Kept (valid on both)', 'Removed (Unknown/Missing)'],
  2328. ncol=2, frameon=False, loc='lower center', bbox_to_anchor=(0.5, 0.08))
  2329. fig.suptitle(f'Per-model DOI pair kept vs removed (n={n_total} each)', fontsize=11)
  2330. fig.text(0.5, 0.015, 'Per-model view: a pair is "removed" for that model if it returned Unknown/Missing on Right or Inverse (independent of the global filter).',
  2331. ha='center', va='bottom', fontsize=6.5)
  2332. plt.show()
  2333. print("✅ Plot 3: Per-model Kept vs Removed Donuts")
  2334. # %% [markdown]
  2335. # ## 4c
  2336. # %%
  2337. # ============================================================
  2338. # Plot 4: Ordinal distributions - Primary Models (Right)
  2339. # ============================================================
  2340. def ordinal_percent_matrix(models, direction, doi_list):
  2341. matrix = np.zeros((len(models), len(VERDICTS)))
  2342. for i, model in enumerate(models):
  2343. counts = np.zeros(len(VERDICTS))
  2344. for doi in doi_list:
  2345. counts[SCORE[lookup_all[(doi, model, direction)]['verdict']]] += 1
  2346. matrix[i] = counts / counts.sum() * 100
  2347. return matrix
  2348. models = PRIMARY_MODELS
  2349. direction = 'Right'
  2350. matrix = ordinal_percent_matrix(models, direction, retained_dois)
  2351. fig, ax = plt.subplots(figsize=(9.2, 5.0))
  2352. y = np.arange(len(models))
  2353. left = np.zeros(len(models))
  2354. for j, (label, color) in enumerate(zip(VERDICT_LABELS, VERDICT_COLORS)):
  2355. ax.barh(y, matrix[:, j], left=left, label=label, color=color)
  2356. for i in range(len(models)):
  2357. value = matrix[i, j]
  2358. if value >= 4.0:
  2359. ax.text(left[i] + value / 2, i, f'{value:.0f}', ha='center', va='center', fontsize=8, color='white')
  2360. left += matrix[:, j]
  2361. ax.set_yticks(y, [DISPLAY[m] for m in models])
  2362. ax.invert_yaxis()
  2363. ax.set_xlim(0, 100)
  2364. ax.set_xlabel('Verdict distribution (%)')
  2365. ax.set_title(f'Normalized ordinal verdict distributions - {direction} hypotheses', pad=34)
  2366. ax.legend(ncol=3, frameon=False, loc='lower center', bbox_to_anchor=(0.5, 1.02))
  2367. ax.grid(axis='x', alpha=0.25)
  2368. footer(fig, f'Right is experimentally supported; Inverse is the DOI-matched incorrect reverse hypothesis. Complete-case n={len(retained_dois)} pairs.')
  2369. plt.show()
  2370. print("✅ Plot 4: Ordinal Distributions - Primary Models (Right)")
  2371. # %%
  2372. # ============================================================
  2373. # Plot 5: Ordinal distributions - Primary Models (Inverse)
  2374. # ============================================================
  2375. models = PRIMARY_MODELS
  2376. direction = 'Inverse'
  2377. matrix = ordinal_percent_matrix(models, direction, retained_dois)
  2378. fig, ax = plt.subplots(figsize=(9.2, 5.0))
  2379. y = np.arange(len(models))
  2380. left = np.zeros(len(models))
  2381. for j, (label, color) in enumerate(zip(VERDICT_LABELS, VERDICT_COLORS)):
  2382. ax.barh(y, matrix[:, j], left=left, label=label, color=color)
  2383. for i in range(len(models)):
  2384. value = matrix[i, j]
  2385. if value >= 4.0:
  2386. ax.text(left[i] + value / 2, i, f'{value:.0f}', ha='center', va='center', fontsize=8, color='white')
  2387. left += matrix[:, j]
  2388. ax.set_yticks(y, [DISPLAY[m] for m in models])
  2389. ax.invert_yaxis()
  2390. ax.set_xlim(0, 100)
  2391. ax.set_xlabel('Verdict distribution (%)')
  2392. ax.set_title(f'Normalized ordinal verdict distributions - {direction} hypotheses', pad=34)
  2393. ax.legend(ncol=3, frameon=False, loc='lower center', bbox_to_anchor=(0.5, 1.02))
  2394. ax.grid(axis='x', alpha=0.25)
  2395. footer(fig, f'Right is experimentally supported; Inverse is the DOI-matched incorrect reverse hypothesis. Complete-case n={len(retained_dois)} pairs.')
  2396. plt.show()
  2397. print("✅ Plot 5: Ordinal Distributions - Primary Models (Inverse)")
  2398. # %% [markdown]
  2399. # ## 4d
  2400. # %%
  2401. # ============================================================
  2402. # Plot 7: Exact pair resolution (Lollipop)
  2403. # ============================================================
  2404. fig, ax = plt.subplots(figsize=(7.0, 4.4))
  2405. y = np.arange(len(PRIMARY_MODELS))
  2406. ax.hlines(y, 0, exact_rates, color='#9ecae1', linewidth=2.5, zorder=1)
  2407. ax.scatter(exact_rates, y, color=BAR_COLOR, s=90, zorder=2)
  2408. for yi, value in zip(y, exact_rates):
  2409. ax.text(value + max(exact_rates + [0.01]) * 0.04 + 0.003, yi, f'{value:.1%}', va='center', fontsize=9)
  2410. ax.set_yticks(y, [DISPLAY[m] for m in PRIMARY_MODELS])
  2411. ax.invert_yaxis()
  2412. ax.set_xlim(0, max(exact_rates + [0.05]) * 1.35 + 0.01)
  2413. ax.set_xlabel('Exact DOI-pair resolution rate')
  2414. ax.set_title('Right=Support and matched Inverse=Non-support')
  2415. ax.grid(axis='x', alpha=0.25)
  2416. footer(fig, 'Same metric as the bar chart, shown as a lollipop for comparison.')
  2417. plt.show()
  2418. print("✅ Plot 7: Exact Pair Resolution Lollipop")
  2419. # %% [markdown]
  2420. # ## 4f
  2421. # %%
  2422. # ============================================================================
  2423. # PLOT 1: Trade-off: edge accuracy vs benchmark score
  2424. # ============================================================================
  2425. import pandas as pd
  2426. import numpy as np
  2427. import matplotlib.pyplot as plt
  2428. import matplotlib as mpl
  2429. import re
  2430. import os
  2431. # Configuration
  2432. PATH_BASE = "judge_results_base_resolved_common.csv"
  2433. PATH_FT3 = "judge_results_3_2_resolved_common.csv"
  2434. PATH_FT4 = "judge_results_4_2_resolved_common.csv"
  2435. PATH_BENCH = "model-accuracy.csv"
  2436. OUTDIR = "outputs"
  2437. os.makedirs(OUTDIR, exist_ok=True)
  2438. MODEL_ORDER = ["Base", "Model 1", "Model 2"]
  2439. MODEL_COLORS = {"Base": "#2C5F9E", "Model 1": "#E8912D", "Model 2": "#3C8C42"}
  2440. MODEL_MARKERS = {"Base": "o", "Model 1": "s", "Model 2": "^"}
  2441. mpl.rcParams.update({
  2442. "svg.fonttype": "none",
  2443. "font.family": "sans-serif",
  2444. "font.sans-serif": ["Arial", "Helvetica", "DejaVu Sans"],
  2445. "font.weight": "bold",
  2446. "axes.labelweight": "bold",
  2447. "axes.titleweight": "bold",
  2448. "axes.linewidth": 1.8,
  2449. "axes.edgecolor": "black",
  2450. "xtick.major.width": 1.6,
  2451. "ytick.major.width": 1.6,
  2452. "xtick.major.size": 6,
  2453. "ytick.major.size": 6,
  2454. "xtick.direction": "in",
  2455. "ytick.direction": "in",
  2456. "xtick.labelsize": 12,
  2457. "ytick.labelsize": 12,
  2458. "axes.labelsize": 13,
  2459. "axes.titlesize": 14,
  2460. "legend.fontsize": 11,
  2461. "legend.edgecolor": "black",
  2462. "legend.framealpha": 1.0,
  2463. })
  2464. def style_axes(ax):
  2465. ax.tick_params(which="both", top=True, right=True, direction="in")
  2466. for spine in ax.spines.values():
  2467. spine.set_linewidth(1.8)
  2468. spine.set_color("black")
  2469. def resolve_prediction(row):
  2470. pred = str(row.get("resolved_answer", "")).upper().strip()
  2471. if pred in {"TRUE", "FALSE"}:
  2472. return pred
  2473. raw = str(row.get("judge_raw_response", "")).upper()
  2474. has_true = bool(re.search(r"\bTRUE\b", raw))
  2475. has_false = bool(re.search(r"\bFALSE\b", raw))
  2476. if has_true and not has_false:
  2477. return "TRUE"
  2478. if has_false and not has_true:
  2479. return "FALSE"
  2480. return "OTHER"
  2481. def load_judge_csv(path):
  2482. df = pd.read_csv(path)
  2483. if df["ground_truth"].dtype != bool:
  2484. df["ground_truth"] = df["ground_truth"].astype(str).str.upper().map({"TRUE": True, "FALSE": False})
  2485. df["prediction"] = df.apply(resolve_prediction, axis=1)
  2486. df["answer_cat"] = df["prediction"].map({"TRUE": "Yes", "FALSE": "No", "OTHER": "Unresolved"})
  2487. return df
  2488. def compute_model_stats(df):
  2489. n = len(df)
  2490. committed = df[df["answer_cat"] != "Unresolved"]
  2491. cols = ["Yes", "No", "Unresolved"]
  2492. conf_counts = df.groupby(["ground_truth", "answer_cat"]).size().unstack(fill_value=0).reindex(index=[True, False], columns=cols, fill_value=0)
  2493. conf_frac = conf_counts.div(conf_counts.sum(axis=1), axis=0)
  2494. resolved_frac = len(committed) / n
  2495. true_committed = committed[committed["ground_truth"] == True]
  2496. false_committed = committed[committed["ground_truth"] == False]
  2497. recall_true = (true_committed["answer_cat"] == "Yes").mean() * 100
  2498. recall_false = (false_committed["answer_cat"] == "No").mean() * 100
  2499. affirm_frac = (committed["answer_cat"] == "Yes").mean() * 100
  2500. edge_accuracy = df["accuracy"].mean()
  2501. prevalence = committed["ground_truth"].mean() * 100
  2502. return {"n": n, "conf_counts": conf_counts, "conf_frac": conf_frac, "resolved_frac": resolved_frac,
  2503. "recall_true": recall_true, "recall_false": recall_false, "affirm_frac": affirm_frac,
  2504. "edge_accuracy": edge_accuracy, "prevalence": prevalence}
  2505. dfs = {"Base": load_judge_csv(PATH_BASE), "Model 1": load_judge_csv(PATH_FT3), "Model 2": load_judge_csv(PATH_FT4)}
  2506. stats = {name: compute_model_stats(df) for name, df in dfs.items()}
  2507. bench_df = pd.read_csv(PATH_BENCH)
  2508. bench_df["Model"] = bench_df["Model"].replace({"Finetuned 3": "Model 1", "Finetuned 4": "Model 2"})
  2509. bench_cols = [c for c in bench_df.columns if c not in ("Model", "Node Loss (U/W)", "Edge Loss (U/W)", "Combined Loss (U/W)")]
  2510. bench_df["avg_bench"] = bench_df[bench_cols].mean(axis=1)
  2511. avg_bench = dict(zip(bench_df["Model"], bench_df["avg_bench"]))
  2512. # Plot
  2513. fig, ax = plt.subplots(figsize=(7.5, 6))
  2514. label_offsets = {"Base": (0.012, 0.001), "Model 1": (-0.10, 0.006), "Model 2": (0.012, 0.003)}
  2515. for name in MODEL_ORDER:
  2516. x = stats[name]["edge_accuracy"]
  2517. y = avg_bench[name]
  2518. ax.scatter(x, y, s=260, marker=MODEL_MARKERS[name], color=MODEL_COLORS[name],
  2519. edgecolor="black", linewidth=1.3, label=name, zorder=3)
  2520. dx, dy = label_offsets[name]
  2521. ax.annotate(name, (x, y), xytext=(x + dx, y + dy), fontsize=13, fontweight="bold", color=MODEL_COLORS[name])
  2522. ax.set_xlabel("Edge accuracy (fraction correct, higher is better \u2192)")
  2523. ax.set_ylabel("Average benchmark score (higher is better \u2191)")
  2524. ax.set_title("Trade-off: edge accuracy vs benchmark score")
  2525. ax.set_xlim(0.5, 1.05)
  2526. style_axes(ax)
  2527. legend = ax.legend(title="Model", loc="lower left", markerscale=0.8)
  2528. legend.get_title().set_fontweight("bold")
  2529. fig.tight_layout()
  2530. fig.savefig(f"{OUTDIR}/fig1_tradeoff_accuracy_vs_benchmark.svg")
  2531. plt.show()
  2532. plt.close(fig)
  2533. # %% [markdown]
  2534. # ## 4g
  2535. # %%
  2536. # ============================================================================
  2537. # PLOT 2: Confusion matrices: ground truth vs model's resolved answer
  2538. # ============================================================================
  2539. import pandas as pd
  2540. import numpy as np
  2541. import matplotlib.pyplot as plt
  2542. import matplotlib as mpl
  2543. import re
  2544. import os
  2545. # Configuration
  2546. PATH_BASE = "judge_results_base_resolved_common.csv"
  2547. PATH_FT3 = "judge_results_3_2_resolved_common.csv"
  2548. PATH_FT4 = "judge_results_4_2_resolved_common.csv"
  2549. OUTDIR = "outputs"
  2550. os.makedirs(OUTDIR, exist_ok=True)
  2551. MODEL_ORDER = ["Base", "Model 1", "Model 2"]
  2552. MODEL_COLORS = {"Base": "#2C5F9E", "Model 1": "#E8912D", "Model 2": "#3C8C42"}
  2553. mpl.rcParams.update({
  2554. "svg.fonttype": "none",
  2555. "font.family": "sans-serif",
  2556. "font.sans-serif": ["Arial", "Helvetica", "DejaVu Sans"],
  2557. "font.weight": "bold",
  2558. "axes.labelweight": "bold",
  2559. "axes.titleweight": "bold",
  2560. "axes.linewidth": 1.8,
  2561. "axes.edgecolor": "black",
  2562. "xtick.major.width": 1.6,
  2563. "ytick.major.width": 1.6,
  2564. "xtick.major.size": 6,
  2565. "ytick.major.size": 6,
  2566. "xtick.direction": "in",
  2567. "ytick.direction": "in",
  2568. "xtick.labelsize": 12,
  2569. "ytick.labelsize": 12,
  2570. "axes.labelsize": 13,
  2571. "axes.titlesize": 14,
  2572. "legend.fontsize": 11,
  2573. "legend.edgecolor": "black",
  2574. "legend.framealpha": 1.0,
  2575. })
  2576. def resolve_prediction(row):
  2577. pred = str(row.get("resolved_answer", "")).upper().strip()
  2578. if pred in {"TRUE", "FALSE"}:
  2579. return pred
  2580. raw = str(row.get("judge_raw_response", "")).upper()
  2581. has_true = bool(re.search(r"\bTRUE\b", raw))
  2582. has_false = bool(re.search(r"\bFALSE\b", raw))
  2583. if has_true and not has_false:
  2584. return "TRUE"
  2585. if has_false and not has_true:
  2586. return "FALSE"
  2587. return "OTHER"
  2588. def load_judge_csv(path):
  2589. df = pd.read_csv(path)
  2590. if df["ground_truth"].dtype != bool:
  2591. df["ground_truth"] = df["ground_truth"].astype(str).str.upper().map({"TRUE": True, "FALSE": False})
  2592. df["prediction"] = df.apply(resolve_prediction, axis=1)
  2593. df["answer_cat"] = df["prediction"].map({"TRUE": "Yes", "FALSE": "No", "OTHER": "Unresolved"})
  2594. return df
  2595. def compute_model_stats(df):
  2596. n = len(df)
  2597. committed = df[df["answer_cat"] != "Unresolved"]
  2598. cols = ["Yes", "No", "Unresolved"]
  2599. conf_counts = df.groupby(["ground_truth", "answer_cat"]).size().unstack(fill_value=0).reindex(index=[True, False], columns=cols, fill_value=0)
  2600. conf_frac = conf_counts.div(conf_counts.sum(axis=1), axis=0)
  2601. resolved_frac = len(committed) / n
  2602. true_committed = committed[committed["ground_truth"] == True]
  2603. false_committed = committed[committed["ground_truth"] == False]
  2604. recall_true = (true_committed["answer_cat"] == "Yes").mean() * 100
  2605. recall_false = (false_committed["answer_cat"] == "No").mean() * 100
  2606. affirm_frac = (committed["answer_cat"] == "Yes").mean() * 100
  2607. edge_accuracy = df["accuracy"].mean()
  2608. prevalence = committed["ground_truth"].mean() * 100
  2609. return {"n": n, "conf_counts": conf_counts, "conf_frac": conf_frac, "resolved_frac": resolved_frac,
  2610. "recall_true": recall_true, "recall_false": recall_false, "affirm_frac": affirm_frac,
  2611. "edge_accuracy": edge_accuracy, "prevalence": prevalence}
  2612. dfs = {"Base": load_judge_csv(PATH_BASE), "Model 1": load_judge_csv(PATH_FT3), "Model 2": load_judge_csv(PATH_FT4)}
  2613. stats = {name: compute_model_stats(df) for name, df in dfs.items()}
  2614. # Plot
  2615. fig, axes = plt.subplots(1, 3, figsize=(15, 5.2))
  2616. row_labels = ["Actually\nTrue", "Actually\nFalse"]
  2617. col_labels = ["Yes\n(True)", "No\n(False)", "Unresolved"]
  2618. im = None
  2619. for ax, name in zip(axes, MODEL_ORDER):
  2620. frac = stats[name]["conf_frac"].values
  2621. counts = stats[name]["conf_counts"].values
  2622. im = ax.imshow(frac, cmap="Blues", vmin=0, vmax=1, aspect="auto")
  2623. for i in range(2):
  2624. for j in range(3):
  2625. val = frac[i, j]
  2626. txt_color = "white" if val > 0.6 else "black"
  2627. ax.text(j, i, f"{counts[i, j]:,}\n{val*100:.1f}%", ha="center", va="center",
  2628. fontsize=12, fontweight="bold", color=txt_color)
  2629. ax.set_xticks(range(3))
  2630. ax.set_xticklabels(col_labels, fontsize=11)
  2631. ax.set_yticks(range(2))
  2632. ax.set_yticklabels(row_labels, fontsize=11)
  2633. ax.set_title(name, color=MODEL_COLORS[name], fontsize=14, fontweight="bold")
  2634. ax.set_xlabel("Model's answer")
  2635. if name == "Base":
  2636. ax.set_ylabel("Ground truth")
  2637. for spine in ax.spines.values():
  2638. spine.set_visible(True)
  2639. spine.set_linewidth(1.8)
  2640. spine.set_color("black")
  2641. ax.tick_params(length=0)
  2642. ax.set_xticks(np.arange(-0.5, 3, 1), minor=True)
  2643. ax.set_yticks(np.arange(-0.5, 2, 1), minor=True)
  2644. ax.grid(which="minor", color="white", linewidth=2)
  2645. ax.tick_params(which="minor", length=0)
  2646. fig.suptitle("Confusion matrices: ground truth vs model's resolved answer", fontsize=15, fontweight="bold")
  2647. cbar = fig.colorbar(im, ax=axes, fraction=0.025, pad=0.02)
  2648. cbar.set_label("Row-normalised fraction", fontsize=11, fontweight="bold")
  2649. fig.savefig(f"{OUTDIR}/fig2_confusion_matrices.svg", bbox_inches="tight")
  2650. plt.show()
  2651. plt.close(fig)
  2652. # %% [markdown]
  2653. # ## 4h
  2654. # %%
  2655. # ============================================================
  2656. # Plot 8: Ordinal distributions - Alternative Models (Right)
  2657. # ============================================================
  2658. models = ALT_MODELS
  2659. direction = 'Right'
  2660. matrix = ordinal_percent_matrix(models, direction, retained_dois)
  2661. fig, ax = plt.subplots(figsize=(9.2, 5.0))
  2662. y = np.arange(len(models))
  2663. left = np.zeros(len(models))
  2664. for j, (label, color) in enumerate(zip(VERDICT_LABELS, VERDICT_COLORS)):
  2665. ax.barh(y, matrix[:, j], left=left, label=label, color=color)
  2666. for i in range(len(models)):
  2667. value = matrix[i, j]
  2668. if value >= 4.0:
  2669. ax.text(left[i] + value / 2, i, f'{value:.0f}', ha='center', va='center', fontsize=8, color='white')
  2670. left += matrix[:, j]
  2671. ax.set_yticks(y, [DISPLAY[m] for m in models])
  2672. ax.invert_yaxis()
  2673. ax.set_xlim(0, 100)
  2674. ax.set_xlabel('Verdict distribution (%)')
  2675. ax.set_title(f'Normalized ordinal verdict distributions - {direction} hypotheses', pad=34)
  2676. ax.legend(ncol=3, frameon=False, loc='lower center', bbox_to_anchor=(0.5, 1.02))
  2677. ax.grid(axis='x', alpha=0.25)
  2678. footer(fig, f'Right is experimentally supported; Inverse is the DOI-matched incorrect reverse hypothesis. Complete-case n={len(retained_dois)} pairs.')
  2679. plt.show()
  2680. print("✅ Plot 8: Ordinal Distributions - Alternative Models (Right)")
  2681. # %%
  2682. # ============================================================
  2683. # Plot 9: Ordinal distributions - Alternative Models (Inverse)
  2684. # ============================================================
  2685. models = ALT_MODELS
  2686. direction = 'Inverse'
  2687. matrix = ordinal_percent_matrix(models, direction, retained_dois)
  2688. fig, ax = plt.subplots(figsize=(9.2, 5.0))
  2689. y = np.arange(len(models))
  2690. left = np.zeros(len(models))
  2691. for j, (label, color) in enumerate(zip(VERDICT_LABELS, VERDICT_COLORS)):
  2692. ax.barh(y, matrix[:, j], left=left, label=label, color=color)
  2693. for i in range(len(models)):
  2694. value = matrix[i, j]
  2695. if value >= 4.0:
  2696. ax.text(left[i] + value / 2, i, f'{value:.0f}', ha='center', va='center', fontsize=8, color='white')
  2697. left += matrix[:, j]
  2698. ax.set_yticks(y, [DISPLAY[m] for m in models])
  2699. ax.invert_yaxis()
  2700. ax.set_xlim(0, 100)
  2701. ax.set_xlabel('Verdict distribution (%)')
  2702. ax.set_title(f'Normalized ordinal verdict distributions - {direction} hypotheses', pad=34)
  2703. ax.legend(ncol=3, frameon=False, loc='lower center', bbox_to_anchor=(0.5, 1.02))
  2704. ax.grid(axis='x', alpha=0.25)
  2705. footer(fig, f'Right is experimentally supported; Inverse is the DOI-matched incorrect reverse hypothesis. Complete-case n={len(retained_dois)} pairs.')
  2706. plt.show()
  2707. print("✅ Plot 9: Ordinal Distributions - Alternative Models (Inverse)")
  2708. # %% [markdown]
  2709. # ## 4i
  2710. # %%
  2711. """
  2712. Single-plot version: verdicts by level, drawn as a lollipop (dot) plot.
  2713. """
  2714. %matplotlib inline
  2715. import ast
  2716. import os
  2717. # Remove matplotlib.use("Agg") - let it use default backend
  2718. import matplotlib.pyplot as plt
  2719. import numpy as np
  2720. import pandas as pd
  2721. # Fix for __file__ not being defined in interactive environments
  2722. try:
  2723. HERE = os.path.dirname(os.path.abspath(__file__))
  2724. except NameError:
  2725. HERE = os.getcwd()
  2726. OUTDIR = os.path.join(HERE, "verdict_analysis")
  2727. # Updated paths for BioChatter and Escargot
  2728. BASE_PATH = "/storage/Arushi/090526_EvoAge/kg_formation/all_figures/FIG_4"
  2729. BIOCHATTER_CSV = os.path.join(BASE_PATH, "biochatter_hypothesis_results_medgemma.csv")
  2730. ESCARGOT_CSV = os.path.join(BASE_PATH, "escargot_hypothesis_results_medgemma.csv")
  2731. EVOAGE_CSV = ("/storage/Arushi/090526_EvoAge/multiagent_hypo/"
  2732. "evoage_100_right_inverse/Right_hypothesis_Extract_entities/"
  2733. "hypothesis_pipeline_results_final_with_percentile_buckets_full.csv")
  2734. LADDER = ["no_support", "weak_support", "partial_support", "support", "strong_support"]
  2735. LABEL = ["No\nsupport", "Weak\nsupport", "Partial\nsupport", "Support", "Strong\nsupport"]
  2736. RANK = {v: i for i, v in enumerate(LADDER)}
  2737. SYSTEMS = ["EvoAge", "Escargot", "BioChatter"]
  2738. COLOR = {"EvoAge": "#2a78d6", "Escargot": "#eb6834", "BioChatter": "#1baf7a"}
  2739. INK, INK_SOFT, GRID, SURFACE = "#0b0b0b", "#52514e", "#e3e2de", "#fcfcfb"
  2740. def parse_evoage_verdict(value):
  2741. try:
  2742. return (ast.literal_eval(value) or {}).get("verdict")
  2743. except (ValueError, SyntaxError, TypeError):
  2744. return None
  2745. def load() -> pd.DataFrame:
  2746. def key(series):
  2747. return series.astype(str).str.strip().str.lower().str.rstrip("/")
  2748. def clean(series):
  2749. cleaned = series.astype(str).str.strip().str.lower()
  2750. return cleaned.where(cleaned.isin(LADDER))
  2751. for csv_path in [EVOAGE_CSV, ESCARGOT_CSV, BIOCHATTER_CSV]:
  2752. if not os.path.exists(csv_path):
  2753. print(f"Warning: File not found: {csv_path}")
  2754. evo = pd.read_csv(EVOAGE_CSV)
  2755. esc = pd.read_csv(ESCARGOT_CSV)
  2756. bio = pd.read_csv(BIOCHATTER_CSV)
  2757. evo = evo.assign(EvoAge=evo["EvoAge_hypothesis_response"].apply(parse_evoage_verdict),
  2758. _key=key(evo["DOI"]))
  2759. esc = esc.assign(Escargot=esc["Medgemma_verdict"], _key=key(esc["DOI"]))
  2760. bio = bio.assign(BioChatter=bio["Medgemma_verdict"], _key=key(bio["DOI"]))
  2761. merged = (evo[["_key", "Title", "DOI", "Right Hypothesis", "EvoAge"]]
  2762. .merge(esc[["_key", "Escargot"]], on="_key", how="inner", validate="1:1")
  2763. .merge(bio[["_key", "BioChatter"]], on="_key", how="inner", validate="1:1"))
  2764. for system in SYSTEMS:
  2765. merged[system] = clean(merged[system])
  2766. return merged.drop(columns="_key")
  2767. def plot_lollipop(frame):
  2768. fig, ax = plt.subplots(figsize=(7.6, 5.0), facecolor=SURFACE)
  2769. ax.set_facecolor(SURFACE)
  2770. counts = {s: frame[s].value_counts().reindex(LADDER).fillna(0).astype(int)
  2771. for s in SYSTEMS}
  2772. y_base = np.arange(len(LADDER))[::-1]
  2773. offsets = {"EvoAge": 0.23, "Escargot": 0.0, "BioChatter": -0.23}
  2774. for system in SYSTEMS:
  2775. y = y_base + offsets[system]
  2776. values = counts[system].values
  2777. ax.hlines(y, 0, values, color=COLOR[system], linewidth=1.8, alpha=0.55)
  2778. ax.plot(values, y, "o", markersize=8, color=COLOR[system],
  2779. markeredgecolor=SURFACE, markeredgewidth=1.4,
  2780. label=system, linestyle="none", zorder=3)
  2781. for yi, value in zip(y, values):
  2782. if value > 0:
  2783. ax.text(value + 2.4, yi, str(value), va="center",
  2784. fontsize=8.5, color=INK_SOFT)
  2785. for side in ("top", "right"):
  2786. ax.spines[side].set_visible(False)
  2787. for side in ("left", "bottom"):
  2788. ax.spines[side].set_color(GRID)
  2789. ax.tick_params(colors=INK_SOFT, labelsize=9, length=3)
  2790. ax.set_axisbelow(True)
  2791. ax.xaxis.grid(True, color=GRID, linewidth=0.7)
  2792. ax.set_yticks(y_base)
  2793. ax.set_yticklabels([l.replace("\n", " ") for l in LABEL], fontsize=10, color=INK)
  2794. ax.set_ylim(-0.6, len(LADDER) - 0.4)
  2795. ax.set_xlim(0, max(108, len(frame) * 1.08))
  2796. ax.set_xlabel("Hypotheses (n)", fontsize=10, color=INK_SOFT)
  2797. ax.set_title(f"Verdicts by level across {len(frame)} biological hypotheses",
  2798. loc="left", fontsize=12, color=INK, pad=12, fontweight="bold")
  2799. ax.legend(frameon=False, fontsize=10, loc="lower right", labelcolor=INK_SOFT)
  2800. fig.tight_layout()
  2801. return fig
  2802. def main() -> None:
  2803. os.makedirs(OUTDIR, exist_ok=True)
  2804. print(f"BioChatter file: {BIOCHATTER_CSV}")
  2805. print(f"Escargot file: {ESCARGOT_CSV}")
  2806. print(f"EvoAge file: {EVOAGE_CSV}")
  2807. frame = load()
  2808. print(f"joined {len(frame)} hypotheses on DOI")
  2809. for system in SYSTEMS:
  2810. n = frame[system].isin(LADDER[1:]).sum()
  2811. print(f" {system:<11} any support: {n}/{len(frame)} ({100*n/len(frame):.0f}%)")
  2812. fig = plot_lollipop(frame)
  2813. svg = os.path.join(OUTDIR, "figure4_verdict_lollipop.svg")
  2814. png = os.path.join(OUTDIR, "figure4_verdict_lollipop.png")
  2815. fig.savefig(svg, format="svg", facecolor=SURFACE, bbox_inches="tight")
  2816. fig.savefig(png, dpi=200, facecolor=SURFACE, bbox_inches="tight")
  2817. plt.show() # This will display the plot
  2818. plt.close(fig)
  2819. print(f"\nwrote {svg}\nwrote {png}")
  2820. if __name__ == "__main__":
  2821. main()
  2822. # %% [markdown]
  2823. # ## Main Figure 5
  2824. # %% [markdown]
  2825. # ## 5b
  2826. # %%
  2827. # This magic command makes plots appear in the notebook
  2828. %matplotlib inline
  2829. """
  2830. log2fc_coherence_plot.py
  2831. Cross-condition coherence plot using log2 fold change (log2FC) instead of
  2832. difference scores.
  2833. For each gene:
  2834. log2FC_c1 = log2(Mean_KO_ratio_c1 / Mean_Control_ratio_c1) # Heat shock 37C vs 30C
  2835. log2FC_c2 = log2(Mean_KO_ratio_c2 / Mean_Control_ratio_c2) # Thermo pulsing vs 30C
  2836. Color = significance status (BH-FDR q <= 0.05), NOT EvoAge verdict:
  2837. grey = not significant
  2838. red = heat shock significant only
  2839. blue = thermo pulsing significant only
  2840. green = significant in both conditions
  2841. Spearman rho computed across all 976 genes (no filtering).
  2842. INPUT : /storage/Arushi/090526_EvoAge/kg_formation/all_figures/yeast-exp/EvoAge_plots/03_log2fc_coherence/input_EvoAge_log2FC_merged_data.csv
  2843. OUTPUT: coherence_log2fc.png / .svg in the same directory
  2844. """
  2845. import warnings
  2846. warnings.filterwarnings('ignore')
  2847. import matplotlib.pyplot as plt
  2848. import numpy as np
  2849. import pandas as pd
  2850. from matplotlib.lines import Line2D
  2851. from scipy import stats
  2852. import os
  2853. from pathlib import Path
  2854. # ── Configuration ──────────────────────────────────────────────────────
  2855. # Input file path
  2856. INPUT_FILE = '/storage/Arushi/090526_EvoAge/kg_formation/all_figures/yeast-exp/EvoAge_plots/03_log2fc_coherence/input_EvoAge_log2FC_merged_data.csv'
  2857. # Output directory (same as input file location)
  2858. OUTPUT_DIR = os.path.dirname(INPUT_FILE)
  2859. OUTPUT_BASE = 'coherence_log2fc'
  2860. # ── Gene name mapping (verified from EvoAge_complete_merged_data.csv) ──
  2861. GENE_NAMES = {
  2862. "YGR036C": "CAX4", "YDL219W": "DTD1", "YDR098C": "GRX3", "YGR155W": "CYS4",
  2863. "YDR226W": "ADK1", "YBR035C": "PDX3", "YBL099W": "ATP1", "YGL012W": "ERG4",
  2864. "YGR087C": "PDC6", "YOR316C": "COT1",
  2865. }
  2866. # Label offsets to prevent text overlap
  2867. LABEL_OFFSETS = {
  2868. "CAX4": (8, 8), "DTD1": (8, -10), "GRX3": (-10, 8), "CYS4": (8, -10), "ADK1": (8, 8),
  2869. "PDX3": (8, 8), "ATP1": (8, -12), "ERG4": (-12, 8), "PDC6": (8, -10), "COT1": (-12, -12),
  2870. }
  2871. # Significance-based colors
  2872. COLORS = {
  2873. "not_sig": "#B0B0B0", # grey
  2874. "heat_only": "#C0392B", # red
  2875. "pulse_only": "#2874A6", # blue
  2876. "both": "#2E8B57", # green
  2877. }
  2878. def load_and_compute(csv_path):
  2879. """Load merged data and compute log2FC if not already present."""
  2880. print(f"Loading data from: {csv_path}")
  2881. df = pd.read_csv(csv_path)
  2882. print(f"Loaded {len(df)} genes")
  2883. # Compute log2FC if the columns don't already exist
  2884. if "log2FC_c1" not in df.columns:
  2885. df["log2FC_c1"] = np.log2(df["Mean_KO_ratio_c1"] / df["Mean_Control_ratio_c1"])
  2886. if "log2FC_c2" not in df.columns:
  2887. df["log2FC_c2"] = np.log2(df["Mean_KO_ratio_c2"] / df["Mean_Control_ratio_c2"])
  2888. # Significance from BH-FDR q-values (these are NOT changed by log2FC)
  2889. df["sig37"] = df["q_value_c1"] <= 0.05
  2890. df["sigPul"] = df["q_value_c2"] <= 0.05
  2891. df["sig_any"] = df["sig37"] | df["sigPul"]
  2892. # Significance category for coloring
  2893. df["color"] = COLORS["not_sig"]
  2894. df.loc[df["sig37"] & ~df["sigPul"], "color"] = COLORS["heat_only"]
  2895. df.loc[df["sigPul"] & ~df["sig37"], "color"] = COLORS["pulse_only"]
  2896. df.loc[df["sig37"] & df["sigPul"], "color"] = COLORS["both"]
  2897. return df.dropna(subset=["log2FC_c1", "log2FC_c2"])
  2898. def make_plot(df, output_dir, output_base, show_plot=True):
  2899. """Create the log2FC cross-condition coherence plot."""
  2900. x = df["log2FC_c1"]
  2901. y = df["log2FC_c2"]
  2902. sig_any = df["sig_any"]
  2903. fig, ax = plt.subplots(figsize=(7.5, 7))
  2904. # All genes (grey majority, colored minority)
  2905. ax.scatter(x, y, c=df["color"], s=30, alpha=0.55,
  2906. edgecolors="none", rasterized=True, zorder=2)
  2907. # Significant genes on top with black outline
  2908. ax.scatter(x[sig_any], y[sig_any], s=90,
  2909. facecolors=df.loc[sig_any, "color"],
  2910. edgecolors="black", linewidths=1.0, zorder=5)
  2911. # Label the 10 significant genes
  2912. for gid, name in GENE_NAMES.items():
  2913. row = df[df["Gene"] == gid]
  2914. if len(row):
  2915. r = row.iloc[0]
  2916. dx, dy = LABEL_OFFSETS.get(name, (8, 8))
  2917. ax.annotate(
  2918. name, (r["log2FC_c1"], r["log2FC_c2"]),
  2919. fontsize=8.5, fontweight="bold",
  2920. ha="left" if dx >= 0 else "right",
  2921. va="bottom" if dy >= 0 else "top",
  2922. xytext=(dx, dy), textcoords="offset points",
  2923. bbox=dict(boxstyle="round,pad=0.15", facecolor="white",
  2924. edgecolor="none", alpha=0.85),
  2925. zorder=6,
  2926. )
  2927. # y=x diagonal + zero lines
  2928. lo, hi = -1.3, 2.0
  2929. ax.plot([lo, hi], [lo, hi], "k--", alpha=0.35, lw=0.8, zorder=1)
  2930. ax.axhline(0, color="#C3C6CA", lw=0.7, zorder=1)
  2931. ax.axvline(0, color="#C3C6CA", lw=0.7, zorder=1)
  2932. # Spearman (all genes, no filtering)
  2933. rho, p = stats.spearmanr(x, y)
  2934. ax.text(
  2935. 0.03, 0.97,
  2936. f"Spearman ρ = {rho:.2f}\nP = {p:.1e}\nn = {len(df)}",
  2937. transform=ax.transAxes, va="top", ha="left", fontsize=10,
  2938. bbox=dict(boxstyle="round", facecolor="white", alpha=0.9), zorder=7,
  2939. )
  2940. ax.set_xlim(lo, hi)
  2941. ax.set_ylim(lo, hi)
  2942. ax.set_xlabel("Heat shock: log₂(KO/control)", fontsize=12)
  2943. ax.set_ylabel("Thermo pulsing: log₂(KO/control)", fontsize=12)
  2944. ax.set_title("Cross-condition coherence (log₂FC, n=976)", fontsize=13, fontweight="bold")
  2945. # Legend
  2946. legend_elems = [
  2947. Line2D([0], [0], marker="o", color="w", markerfacecolor=COLORS["not_sig"], markersize=9, label="Not significant"),
  2948. Line2D([0], [0], marker="o", color="w", markerfacecolor=COLORS["heat_only"], markersize=9, label="Heat shock significant"),
  2949. Line2D([0], [0], marker="o", color="w", markerfacecolor=COLORS["pulse_only"], markersize=9, label="Thermo pulsing significant"),
  2950. Line2D([0], [0], marker="o", color="w", markerfacecolor=COLORS["both"], markersize=9, label="Both significant"),
  2951. ]
  2952. ax.legend(handles=legend_elems, loc="lower right", fontsize=8,
  2953. frameon=True, edgecolor="grey")
  2954. ax.spines["top"].set_visible(False)
  2955. ax.spines["right"].set_visible(False)
  2956. plt.tight_layout()
  2957. # Show plot in notebook
  2958. if show_plot:
  2959. plt.show()
  2960. # Save figures
  2961. for ext in ['png', 'svg']:
  2962. output_path = os.path.join(output_dir, f"{output_base}.{ext}")
  2963. plt.savefig(output_path, dpi=200, bbox_inches="tight")
  2964. print(f"✅ Saved: {output_path}")
  2965. plt.close()
  2966. print(f"\n📊 Statistics:")
  2967. print(f"Spearman ρ = {rho:.5f}, P = {p:.2e}")
  2968. print(f"Genes: n={len(df)}")
  2969. print(f" Not significant: {(~sig_any).sum()}")
  2970. print(f" Heat shock only: {(df['sig37'] & ~df['sigPul']).sum()}")
  2971. print(f" Thermo pulsing only: {(df['sigPul'] & ~df['sig37']).sum()}")
  2972. print(f" Both significant: {(df['sig37'] & df['sigPul']).sum()}")
  2973. # ── Main execution ──────────────────────────────────────────────────────
  2974. # Load data
  2975. df = load_and_compute(INPUT_FILE)
  2976. # Generate plot (will show in notebook and save to files)
  2977. make_plot(df, OUTPUT_DIR, OUTPUT_BASE, show_plot=True)
  2978. print(f"\n✅ All done! Plots saved in: {OUTPUT_DIR}")
  2979. # %% [markdown]
  2980. # ## 5c
  2981. # %%
  2982. # ============================================
  2983. # PLOT 1: Heat Shock (37 °C) vs 30 °C
  2984. # ============================================
  2985. import numpy as np
  2986. import pandas as pd
  2987. import matplotlib.pyplot as plt
  2988. from matplotlib.lines import Line2D
  2989. import matplotlib.ticker as ticker
  2990. import os
  2991. import warnings
  2992. # Suppress font warnings
  2993. warnings.filterwarnings('ignore', category=UserWarning, module='matplotlib')
  2994. # Set font to a standard one that exists in most systems
  2995. plt.rcParams['font.family'] = 'sans-serif'
  2996. plt.rcParams['font.sans-serif'] = ['Arial', 'DejaVu Sans', 'Helvetica', 'sans-serif']
  2997. # Enable inline plotting in Jupyter
  2998. %matplotlib inline
  2999. # Define colors and labels
  3000. VERDICT_COLORS = {
  3001. 'support': '#6FAE84',
  3002. 'weak_support': '#7E7AAE',
  3003. 'partial_support': '#E0D55A',
  3004. 'no_support': '#D9754A',
  3005. }
  3006. VERDICT_LABELS = {
  3007. 'support': 'Support',
  3008. 'weak_support': 'Weak support',
  3009. 'partial_support': 'Partial support',
  3010. 'no_support': 'No support',
  3011. }
  3012. def to_bool(s):
  3013. return str(s).strip().lower() in {'true','1','yes'}
  3014. def prep(df):
  3015. out = df.copy()
  3016. out['Sig_37C'] = out['Sig_37C'].map(to_bool)
  3017. out['Sig_Pulser'] = out['Sig_Pulser'].map(to_bool)
  3018. out['neglog10_q_37C'] = -np.log10(pd.to_numeric(out['q_value_c1'], errors='coerce').clip(lower=1e-300))
  3019. out['neglog10_q_Pulser'] = -np.log10(pd.to_numeric(out['q_value_c2'], errors='coerce').clip(lower=1e-300))
  3020. out['Score_c1'] = pd.to_numeric(out['Score_c1'], errors='coerce')
  3021. out['Score_c2'] = pd.to_numeric(out['Score_c2'], errors='coerce')
  3022. # Use standard names like ERG4, ATP1; fall back to systematic names if needed
  3023. out['label_name'] = out['Standard_Name'].fillna('').astype(str)
  3024. out.loc[out['label_name'].str.strip()=='', 'label_name'] = out.loc[out['label_name'].str.strip()=='', 'Gene']
  3025. return out
  3026. def style_axis_heatshock(ax):
  3027. """Style the axis for heat shock plot"""
  3028. ax.set_title('Heat shock (37 °C) vs 30 °C', fontsize=12, pad=8)
  3029. ax.text(0.5, 0.985, 'n = 976 genes; all FDR values shown', transform=ax.transAxes,
  3030. ha='center', va='top', fontsize=8)
  3031. ax.set_xlabel('Thermal-protection score', fontsize=11)
  3032. ax.set_ylabel(r'-log$_{10}$(BH FDR)', fontsize=11)
  3033. ax.axvline(0, color='#D0D0D0', lw=1.2, zorder=0)
  3034. qline = -np.log10(0.05)
  3035. ax.axhline(qline, color='#9A9A9A', lw=1.6, ls=(0,(4,3)), zorder=0)
  3036. ax.text(0.985, qline+0.05, 'dashed line: q = 0.05', transform=ax.get_yaxis_transform(),
  3037. ha='right', va='bottom', fontsize=8, color='#555555')
  3038. ax.set_xlim(-0.7, 2.6)
  3039. ax.set_xticks([-0.5, 0, 0.5, 1, 1.5, 2, 2.5])
  3040. ax.set_ylim(0, 110)
  3041. ax.set_yscale('symlog', linthresh=1.0, linscale=0.85, base=10)
  3042. ax.set_yticks([0, 1, 2, 5, 10, 20, 50, 100])
  3043. ax.get_yaxis().set_major_formatter(ticker.ScalarFormatter())
  3044. ax.tick_params(width=1.0, length=4, labelsize=9)
  3045. ax.spines['top'].set_visible(False)
  3046. ax.spines['right'].set_visible(False)
  3047. for s in ['left','bottom']:
  3048. ax.spines[s].set_linewidth(1.1)
  3049. def make_handles():
  3050. handles = [Line2D([0],[0], marker='o', linestyle='None', markersize=4,
  3051. markerfacecolor='#C9C9C9', markeredgecolor='none', label='Not significant')]
  3052. for key in ['support','weak_support','partial_support','no_support']:
  3053. handles.append(Line2D([0],[0], marker='o', linestyle='None', markersize=5,
  3054. markerfacecolor=VERDICT_COLORS[key], markeredgecolor='none',
  3055. label=VERDICT_LABELS[key]))
  3056. return handles
  3057. def label_points_heatshock(ax, subdf):
  3058. """Label points for heat shock plot"""
  3059. offsets = {
  3060. 'YGR036C': (-0.04, 0.0, 'right'), # CAX4
  3061. 'YDL219W': (0.04, 0.0, 'left'), # DTD1
  3062. 'YDR098C': (0.04, 0.0, 'left'), # GRX3
  3063. 'YGR155W': (0.04, 0.0, 'left'), # CYS4
  3064. 'YDR226W': (0.04, -0.08, 'left'), # ADK1
  3065. 'YBR035C': (0.04, -0.06, 'left'), # PDX3
  3066. 'YBL099W': (0.04, 0.0, 'left'), # ATP1
  3067. 'YGL012W': (0.04, 0.02, 'left'), # ERG4
  3068. 'YGR087C': (0.04, -0.06, 'left'), # PDC6
  3069. 'YOR316C': (-0.05, 0.02, 'right'), # COT1
  3070. }
  3071. for _, r in subdf.iterrows():
  3072. dx, dy, ha = offsets.get(r['Gene'], (0.04, 0.0, 'left'))
  3073. ax.text(r['Score_c1']+dx, r['neglog10_q_37C']+dy, r['label_name'],
  3074. fontsize=8.8, ha=ha, va='center', color='#333333')
  3075. # ============================================
  3076. # LOAD DATA AND CREATE HEAT SHOCK PLOT
  3077. # ============================================
  3078. # Update this path to your CSV file
  3079. input_file = 'EvoAge_effect_FDR_input.csv'
  3080. if not os.path.exists(input_file):
  3081. print(f"⚠️ Error: File '{input_file}' not found!")
  3082. print(f"Current directory: {os.getcwd()}")
  3083. else:
  3084. # Load and prepare data
  3085. print(f"✅ Loading data from: {input_file}")
  3086. df = prep(pd.read_csv(input_file))
  3087. print(f"✅ Data loaded. Shape: {df.shape}")
  3088. print(f"✅ Significant genes at 37°C: {df['Sig_37C'].sum()}")
  3089. # Create the plot
  3090. fig, ax = plt.subplots(figsize=(4.1,4.0), constrained_layout=True)
  3091. # Style the axis
  3092. style_axis_heatshock(ax)
  3093. # Plot all points in grey
  3094. ax.scatter(df['Score_c1'], df['neglog10_q_37C'], s=10, color='#C9C9C9',
  3095. alpha=0.85, edgecolors='none', zorder=1)
  3096. # Plot significant points with colors based on verdict
  3097. sig = df[df['Sig_37C']].copy()
  3098. for key in ['support','weak_support','partial_support','no_support']:
  3099. ss = sig[sig['verdict'] == key]
  3100. if len(ss):
  3101. ax.scatter(ss['Score_c1'], ss['neglog10_q_37C'], s=24,
  3102. color=VERDICT_COLORS[key], edgecolors='none', zorder=3)
  3103. # Add labels
  3104. label_points_heatshock(ax, sig)
  3105. # Add legend
  3106. ax.legend(handles=make_handles(), title='EvoAge verdict', loc='lower right',
  3107. frameon=False, fontsize=8, title_fontsize=8.5,
  3108. handletextpad=0.4, borderpad=0.2, labelspacing=0.3)
  3109. # Add panel label
  3110. ax.text(-0.13, 1.04, 'c', transform=ax.transAxes, fontsize=16, fontweight='bold')
  3111. # Display the plot
  3112. plt.show()
  3113. # Optional: Save the figure (uncomment to use)
  3114. # fig.savefig('heatshock_panel.png', bbox_inches='tight', dpi=300)
  3115. # fig.savefig('heatshock_panel.pdf', bbox_inches='tight')
  3116. # print("✅ Figure saved!")
  3117. plt.close(fig)
  3118. # %% [markdown]
  3119. # ## 5d
  3120. # %%
  3121. # ============================================
  3122. # PLOT 2: Thermo Pulsing vs 30 °C
  3123. # ============================================
  3124. import numpy as np
  3125. import pandas as pd
  3126. import matplotlib.pyplot as plt
  3127. from matplotlib.lines import Line2D
  3128. import matplotlib.ticker as ticker
  3129. import os
  3130. import warnings
  3131. # Suppress font warnings
  3132. warnings.filterwarnings('ignore', category=UserWarning, module='matplotlib')
  3133. # Set font to a standard one that exists in most systems
  3134. plt.rcParams['font.family'] = 'sans-serif'
  3135. plt.rcParams['font.sans-serif'] = ['Arial', 'DejaVu Sans', 'Helvetica', 'sans-serif']
  3136. # Enable inline plotting in Jupyter
  3137. %matplotlib inline
  3138. # Define colors and labels
  3139. VERDICT_COLORS = {
  3140. 'support': '#6FAE84',
  3141. 'weak_support': '#7E7AAE',
  3142. 'partial_support': '#E0D55A',
  3143. 'no_support': '#D9754A',
  3144. }
  3145. VERDICT_LABELS = {
  3146. 'support': 'Support',
  3147. 'weak_support': 'Weak support',
  3148. 'partial_support': 'Partial support',
  3149. 'no_support': 'No support',
  3150. }
  3151. def to_bool(s):
  3152. return str(s).strip().lower() in {'true','1','yes'}
  3153. def prep(df):
  3154. out = df.copy()
  3155. out['Sig_37C'] = out['Sig_37C'].map(to_bool)
  3156. out['Sig_Pulser'] = out['Sig_Pulser'].map(to_bool)
  3157. out['neglog10_q_37C'] = -np.log10(pd.to_numeric(out['q_value_c1'], errors='coerce').clip(lower=1e-300))
  3158. out['neglog10_q_Pulser'] = -np.log10(pd.to_numeric(out['q_value_c2'], errors='coerce').clip(lower=1e-300))
  3159. out['Score_c1'] = pd.to_numeric(out['Score_c1'], errors='coerce')
  3160. out['Score_c2'] = pd.to_numeric(out['Score_c2'], errors='coerce')
  3161. # Use standard names like ERG4, ATP1; fall back to systematic names if needed
  3162. out['label_name'] = out['Standard_Name'].fillna('').astype(str)
  3163. out.loc[out['label_name'].str.strip()=='', 'label_name'] = out.loc[out['label_name'].str.strip()=='', 'Gene']
  3164. return out
  3165. def style_axis_thermopulsing(ax):
  3166. """Style the axis for thermo pulsing plot"""
  3167. ax.set_title('Thermo pulsing vs 30 °C', fontsize=12, pad=8)
  3168. ax.text(0.5, 0.985, 'n = 976 genes; all FDR values shown', transform=ax.transAxes,
  3169. ha='center', va='top', fontsize=8)
  3170. ax.set_xlabel('Thermal-protection score', fontsize=11)
  3171. ax.set_ylabel(r'-log$_{10}$(BH FDR)', fontsize=11)
  3172. ax.axvline(0, color='#D0D0D0', lw=1.2, zorder=0)
  3173. qline = -np.log10(0.05)
  3174. ax.axhline(qline, color='#9A9A9A', lw=1.6, ls=(0,(4,3)), zorder=0)
  3175. ax.text(0.985, qline+0.05, 'dashed line: q = 0.05', transform=ax.get_yaxis_transform(),
  3176. ha='right', va='bottom', fontsize=8, color='#555555')
  3177. ax.set_xlim(-0.7, 2.6)
  3178. ax.set_xticks([-0.5, 0, 0.5, 1, 1.5, 2, 2.5])
  3179. ax.set_ylim(0, 110)
  3180. ax.set_yscale('symlog', linthresh=1.0, linscale=0.85, base=10)
  3181. ax.set_yticks([0, 1, 2, 5, 10, 20, 50, 100])
  3182. ax.get_yaxis().set_major_formatter(ticker.ScalarFormatter())
  3183. ax.tick_params(width=1.0, length=4, labelsize=9)
  3184. ax.spines['top'].set_visible(False)
  3185. ax.spines['right'].set_visible(False)
  3186. for s in ['left','bottom']:
  3187. ax.spines[s].set_linewidth(1.1)
  3188. def make_handles():
  3189. handles = [Line2D([0],[0], marker='o', linestyle='None', markersize=4,
  3190. markerfacecolor='#C9C9C9', markeredgecolor='none', label='Not significant')]
  3191. for key in ['support','weak_support','partial_support','no_support']:
  3192. handles.append(Line2D([0],[0], marker='o', linestyle='None', markersize=5,
  3193. markerfacecolor=VERDICT_COLORS[key], markeredgecolor='none',
  3194. label=VERDICT_LABELS[key]))
  3195. return handles
  3196. def label_points_thermopulsing(ax, subdf):
  3197. """Label points for thermo pulsing plot"""
  3198. offsets = {
  3199. 'YGR036C': (-0.04, 0.0, 'right'), # CAX4
  3200. 'YDL219W': (0.04, 0.0, 'left'), # DTD1
  3201. 'YDR098C': (0.04, 0.0, 'left'), # GRX3
  3202. 'YGR155W': (0.04, 0.0, 'left'), # CYS4
  3203. 'YDR226W': (0.04, -0.08, 'left'), # ADK1
  3204. 'YBR035C': (0.04, -0.06, 'left'), # PDX3
  3205. 'YBL099W': (0.04, 0.0, 'left'), # ATP1
  3206. 'YGL012W': (0.04, 0.02, 'left'), # ERG4
  3207. 'YGR087C': (0.04, -0.06, 'left'), # PDC6
  3208. 'YOR316C': (-0.05, 0.02, 'right'), # COT1
  3209. }
  3210. for _, r in subdf.iterrows():
  3211. dx, dy, ha = offsets.get(r['Gene'], (0.04, 0.0, 'left'))
  3212. ax.text(r['Score_c2']+dx, r['neglog10_q_Pulser']+dy, r['label_name'],
  3213. fontsize=8.8, ha=ha, va='center', color='#333333')
  3214. # ============================================
  3215. # LOAD DATA AND CREATE THERMO PULSING PLOT
  3216. # ============================================
  3217. # Update this path to your CSV file
  3218. input_file = 'EvoAge_effect_FDR_input.csv'
  3219. if not os.path.exists(input_file):
  3220. print(f"⚠️ Error: File '{input_file}' not found!")
  3221. print(f"Current directory: {os.getcwd()}")
  3222. else:
  3223. # Load and prepare data
  3224. print(f"✅ Loading data from: {input_file}")
  3225. df = prep(pd.read_csv(input_file))
  3226. print(f"✅ Data loaded. Shape: {df.shape}")
  3227. print(f"✅ Significant genes in pulser: {df['Sig_Pulser'].sum()}")
  3228. # Create the plot
  3229. fig, ax = plt.subplots(figsize=(4.1,4.0), constrained_layout=True)
  3230. # Style the axis
  3231. style_axis_thermopulsing(ax)
  3232. # Plot all points in grey
  3233. ax.scatter(df['Score_c2'], df['neglog10_q_Pulser'], s=10, color='#C9C9C9',
  3234. alpha=0.85, edgecolors='none', zorder=1)
  3235. # Plot significant points with colors based on verdict
  3236. sig = df[df['Sig_Pulser']].copy()
  3237. for key in ['support','weak_support','partial_support','no_support']:
  3238. ss = sig[sig['verdict'] == key]
  3239. if len(ss):
  3240. ax.scatter(ss['Score_c2'], ss['neglog10_q_Pulser'], s=24,
  3241. color=VERDICT_COLORS[key], edgecolors='none', zorder=3)
  3242. # Add labels
  3243. label_points_thermopulsing(ax, sig)
  3244. # Add legend
  3245. ax.legend(handles=make_handles(), title='EvoAge verdict', loc='lower right',
  3246. frameon=False, fontsize=8, title_fontsize=8.5,
  3247. handletextpad=0.4, borderpad=0.2, labelspacing=0.3)
  3248. # Add panel label
  3249. ax.text(-0.13, 1.04, 'd', transform=ax.transAxes, fontsize=16, fontweight='bold')
  3250. # Display the plot
  3251. plt.show()
  3252. # Optional: Save the figure (uncomment to use)
  3253. # fig.savefig('thermopulsing_panel.png', bbox_inches='tight', dpi=300)
  3254. # fig.savefig('thermopulsing_panel.pdf', bbox_inches='tight')
  3255. # print("✅ Figure saved!")
  3256. plt.close(fig)
  3257. # %% [markdown]
  3258. # ## 5e
  3259. # %%
  3260. # This magic command makes plots appear in the notebook
  3261. %matplotlib inline
  3262. """
  3263. EvoAge Overlap Venn Diagram
  3264. ---------------------------
  3265. Creates a Venn diagram showing overlap of significant genes between conditions.
  3266. INPUT : /storage/Arushi/090526_EvoAge/kg_formation/all_figures/yeast-exp/EvoAge_plots/overlap/EvoAge_complete_merged_data.csv
  3267. OUTPUT: 03_overlap_venn.svg / .png in the same directory
  3268. """
  3269. import warnings
  3270. warnings.filterwarnings('ignore')
  3271. from pathlib import Path
  3272. import pandas as pd
  3273. import numpy as np
  3274. import matplotlib.pyplot as plt
  3275. from matplotlib.patches import Circle
  3276. # ── Configuration ──────────────────────────────────────────────────────
  3277. # Input directory and file
  3278. INPUT_DIR = '/storage/Arushi/090526_EvoAge/kg_formation/all_figures/yeast-exp/EvoAge_plots/overlap'
  3279. INPUT_FILE = Path(INPUT_DIR) / 'EvoAge_complete_merged_data.csv'
  3280. OUTPUT_DIR = INPUT_DIR
  3281. # ── Constants (from common.py) ──────────────────────────────────────
  3282. VERDICT_MAP = {
  3283. 'support': 'Support',
  3284. 'weak_support': 'Weak support',
  3285. 'partial_support': 'Partial support',
  3286. 'no_support': 'No support',
  3287. }
  3288. VERDICT_ORDER = ['Support', 'Weak support', 'Partial support', 'No support', 'Unresolved']
  3289. COLORS = {
  3290. 'Support': '#5B4B8A',
  3291. 'Weak support': '#6FAE9B',
  3292. 'Partial support': '#D8B65C',
  3293. 'No support': '#C97B63',
  3294. 'Unresolved': '#CFCFD4',
  3295. '37C': '#B66A5B',
  3296. 'Pulser': '#6E9D8C',
  3297. 'neutral': '#C9CDD2',
  3298. }
  3299. def configure():
  3300. plt.rcParams.update({
  3301. 'font.family': 'DejaVu Sans',
  3302. 'font.size': 9,
  3303. 'axes.titlesize': 11,
  3304. 'axes.labelsize': 10,
  3305. 'xtick.labelsize': 8.5,
  3306. 'ytick.labelsize': 8.5,
  3307. 'legend.fontsize': 8,
  3308. 'axes.linewidth': 0.8,
  3309. 'svg.fonttype': 'none',
  3310. 'pdf.fonttype': 42,
  3311. 'ps.fonttype': 42,
  3312. })
  3313. def to_bool(x):
  3314. return str(x).strip().lower() in {'true', '1', 'yes'}
  3315. def load_data(path):
  3316. """Load and prepare data for Venn diagram."""
  3317. print(f"Loading data from: {path}")
  3318. df = pd.read_csv(path)
  3319. print(f"Loaded {len(df)} genes")
  3320. # Create boolean columns for significance
  3321. df['sig_37C'] = df['Sig_37C'].apply(to_bool)
  3322. df['sig_Pulser'] = df['Sig_Pulser'].apply(to_bool)
  3323. df['sig_both'] = df['sig_37C'] & df['sig_Pulser']
  3324. # Get display names
  3325. df['display_name'] = df['Standard_Name'].fillna(df['Gene'])
  3326. return df
  3327. def create_venn(df, output_dir, show_plot=True):
  3328. """Create Venn diagram showing overlap of significant genes."""
  3329. # Get gene lists for each category
  3330. only37 = df[df['sig_37C'] & ~df['sig_Pulser']]['display_name'].tolist()
  3331. both = df[df['sig_both']]['display_name'].tolist()
  3332. onlyp = df[df['sig_Pulser'] & ~df['sig_37C']]['display_name'].tolist()
  3333. # Print statistics
  3334. print(f"\nSignificant genes:")
  3335. print(f" 37°C Heat shock only: {len(only37)} genes")
  3336. print(f" Thermo pulsing only: {len(onlyp)} genes")
  3337. print(f" Both conditions: {len(both)} genes")
  3338. print(f" Total unique significant genes: {len(only37) + len(both) + len(onlyp)}")
  3339. # Create figure
  3340. fig, ax = plt.subplots(figsize=(6.2, 5))
  3341. ax.set_aspect('equal')
  3342. ax.axis('off')
  3343. # Draw circles
  3344. ax.add_patch(Circle((.43, .62), .28,
  3345. facecolor=COLORS['37C'],
  3346. edgecolor=COLORS['37C'],
  3347. alpha=.35, lw=1.5))
  3348. ax.add_patch(Circle((.65, .62), .28,
  3349. facecolor=COLORS['Pulser'],
  3350. edgecolor=COLORS['Pulser'],
  3351. alpha=.35, lw=1.5))
  3352. # Add numbers
  3353. ax.text(.27, .64, len(only37), ha='center', va='center',
  3354. fontsize=18, fontweight='bold')
  3355. ax.text(.54, .64, len(both), ha='center', va='center',
  3356. fontsize=18, fontweight='bold')
  3357. ax.text(.81, .64, len(onlyp), ha='center', va='center',
  3358. fontsize=18, fontweight='bold')
  3359. # Add labels
  3360. ax.text(.23, .94, 'Heat shock (37 °C)', ha='center', fontweight='bold')
  3361. ax.text(.84, .94, 'Thermal pulsing', ha='center', fontweight='bold')
  3362. # Add gene lists (truncated if too long)
  3363. max_genes_display = 10
  3364. gene_lists = []
  3365. for gene_list in [only37, both, onlyp]:
  3366. if len(gene_list) > max_genes_display:
  3367. gene_lists.append(' · '.join(gene_list[:max_genes_display]) + f'\n... and {len(gene_list) - max_genes_display} more')
  3368. elif len(gene_list) > 0:
  3369. gene_lists.append(' · '.join(gene_list))
  3370. else:
  3371. gene_lists.append('(none)')
  3372. ax.text(.18, .22, '37 °C only\n' + gene_lists[0],
  3373. ha='center', va='top', fontsize=8)
  3374. ax.text(.54, .22, 'Both\n' + gene_lists[1],
  3375. ha='center', va='top', fontsize=8)
  3376. ax.text(.88, .22, 'Pulser only\n' + gene_lists[2],
  3377. ha='center', va='top', fontsize=8)
  3378. # Add statistics
  3379. union = len(only37) + len(both) + len(onlyp)
  3380. jaccard = len(both) / union if union > 0 else 0
  3381. ax.text(.5, .04, f'Union = {union}; Jaccard index = {jaccard:.2f}',
  3382. ha='center', fontsize=9)
  3383. ax.set_xlim(0, 1.05)
  3384. ax.set_ylim(0, 1)
  3385. ax.set_title('Overlap of FDR-significant genes', fontweight='bold')
  3386. plt.tight_layout()
  3387. # Show in notebook
  3388. if show_plot:
  3389. plt.show()
  3390. # Save figures
  3391. output_stem = Path(output_dir) / '03_overlap_venn'
  3392. for ext in ['svg', 'png']:
  3393. output_path = output_stem.with_suffix(f'.{ext}')
  3394. plt.savefig(output_path, dpi=300, bbox_inches='tight')
  3395. print(f"✅ Saved: {output_path}")
  3396. plt.close()
  3397. return fig
  3398. # ── Main execution ──────────────────────────────────────────────────────
  3399. # Setup
  3400. configure()
  3401. # Load data
  3402. df = load_data(INPUT_FILE)
  3403. # Create Venn diagram
  3404. fig = create_venn(df, OUTPUT_DIR, show_plot=True)
  3405. print(f"\n✅ All done! Venn diagram saved in: {OUTPUT_DIR}")
  3406. # %% [markdown]
  3407. # ## 5f
  3408. # %%
  3409. # This magic command makes plots appear in the notebook
  3410. %matplotlib inline
  3411. """
  3412. 01_boxplot_verdict_vs_effectsize
  3413. --------------------------------
  3414. Broken-axis violin + box plot: EvoAge verdict vs max signed thermal-protection score.
  3415. INPUT : /storage/Arushi/090526_EvoAge/kg_formation/all_figures/yeast-exp/EvoAge_plots/01_boxplot_verdict_vs_effectsize/input_boxplot_verdict_data.csv
  3416. OUTPUT: boxplot_verdict_vs_effectsize_brokenaxis.svg / .png / .pdf
  3417. Y-axis: ms = score from the condition (37C or Pulser vs 30C) with larger |value|,
  3418. sign kept. Positive = protective; negative = sensitising.
  3419. Axis break at 1.6 isolates CAX4 (+2.55) and DTD1 (+2.47) from the main distribution.
  3420. """
  3421. import warnings
  3422. warnings.filterwarnings('ignore')
  3423. import pandas as pd
  3424. import numpy as np
  3425. import os
  3426. import matplotlib.pyplot as plt
  3427. from matplotlib.lines import Line2D
  3428. # ── Configuration ──────────────────────────────────────────────────────
  3429. # Input file path
  3430. INPUT_FILE = '/storage/Arushi/090526_EvoAge/kg_formation/all_figures/yeast-exp/EvoAge_plots/01_boxplot_verdict_vs_effectsize/input_boxplot_verdict_data.csv'
  3431. # Output directory (same as input file location)
  3432. OUTPUT_DIR = os.path.dirname(INPUT_FILE)
  3433. DPI = 600
  3434. # ── Load data ──────────────────────────────────────────────────────────
  3435. print(f"Loading data from: {INPUT_FILE}")
  3436. df = pd.read_csv(INPUT_FILE)
  3437. print(f"Loaded {len(df)} genes")
  3438. # ── Prepare data ──────────────────────────────────────────────────────
  3439. groups = ['support', 'partial_support', 'weak_support', 'no_support']
  3440. data = [df[df['vg'] == g]['ms'].dropna().values for g in groups]
  3441. n = [len(d) for d in data]
  3442. labels = ['Support', 'Partial\nsupport', 'Weak\nsupport', 'No\nsupport']
  3443. colors = ['#1b5e20', '#66bb6a', '#fff9c4', '#d32f2f']
  3444. # Significant genes mapping
  3445. sig_genes = {'YGR036C': 'CAX4', 'YDL219W': 'DTD1', 'YDR098C': 'GRX3',
  3446. 'YGR155W': 'CYS4', 'YDR226W': 'ADK1', 'YBR035C': 'PDX3',
  3447. 'YBL099W': 'ATP1', 'YGL012W': 'ERG4', 'YOR316C': 'COT1',
  3448. 'YGR087C': 'PDC6'}
  3449. # Get positions for significant genes
  3450. key_pos = {}
  3451. for gene, nm in sig_genes.items():
  3452. r = df[df['Gene'] == gene]
  3453. if len(r) > 0:
  3454. r = r.iloc[0]
  3455. key_pos[nm] = (groups.index(r['vg']) if r['vg'] in groups else 0, r['ms'])
  3456. BREAK = 1.6
  3457. # ── Create figure ──────────────────────────────────────────────────────
  3458. fig = plt.figure(figsize=(9.2, 8.4))
  3459. gs = fig.add_gridspec(2, 1, height_ratios=[1.1, 4.2], hspace=0.06)
  3460. # ── TOP PANEL: extreme outliers (> BREAK) ────────────────────────────────
  3461. ax_top = fig.add_subplot(gs[0])
  3462. for i, (d, color) in enumerate(zip(data, colors)):
  3463. hi = d[d > BREAK]
  3464. if len(hi):
  3465. ax_top.scatter([i] * len(hi), hi, s=60, c=color, edgecolors='black',
  3466. linewidth=1.2, zorder=5)
  3467. # Label CAX4 and DTD1 in top panel
  3468. ax_top.annotate('CAX4', key_pos['CAX4'],
  3469. xytext=(key_pos['CAX4'][0] - 0.25, key_pos['CAX4'][1]),
  3470. fontsize=10, fontweight='bold', ha='right', va='center',
  3471. bbox=dict(boxstyle='round,pad=0.2', facecolor='white',
  3472. edgecolor='#999', linewidth=0.6, alpha=0.95), zorder=10)
  3473. ax_top.annotate('DTD1', key_pos['DTD1'],
  3474. xytext=(key_pos['DTD1'][0] + 0.25, key_pos['DTD1'][1] - 0.08),
  3475. fontsize=10, fontweight='bold', ha='left', va='center',
  3476. bbox=dict(boxstyle='round,pad=0.2', facecolor='white',
  3477. edgecolor='#999', linewidth=0.6, alpha=0.95), zorder=10)
  3478. ax_top.set_ylim(BREAK - 0.05, 2.8)
  3479. ax_top.set_xlim(-1.2, 4.5)
  3480. ax_top.set_xticks([])
  3481. ax_top.spines['top'].set_visible(False)
  3482. ax_top.spines['right'].set_visible(False)
  3483. ax_top.spines['bottom'].set_visible(False)
  3484. ax_top.tick_params(labelsize=9)
  3485. ax_top.yaxis.set_major_locator(plt.MultipleLocator(0.5))
  3486. # ── BOTTOM PANEL: main distribution (<= BREAK) ───────────────────────────
  3487. ax = fig.add_subplot(gs[1])
  3488. # Violin plots
  3489. for i, (d, color) in enumerate(zip(data, colors)):
  3490. d_below = d[d <= BREAK]
  3491. if len(d_below) > 0:
  3492. vp = ax.violinplot(d_below, positions=[i], showextrema=False, widths=0.75)
  3493. for body in vp['bodies']:
  3494. body.set_facecolor(color)
  3495. body.set_alpha(0.28)
  3496. body.set_edgecolor('none')
  3497. # Box plots
  3498. bp = ax.boxplot([d[d <= BREAK] for d in data], positions=[0, 1, 2, 3], widths=0.28,
  3499. patch_artist=True, medianprops=dict(color='black', linewidth=2.2),
  3500. flierprops=dict(marker='o', markerfacecolor='#616161', markersize=4,
  3501. markeredgecolor='#424242', markeredgewidth=0.6, alpha=0.8),
  3502. whiskerprops=dict(linewidth=1.4), capprops=dict(linewidth=1.4))
  3503. for patch, color in zip(bp['boxes'], colors):
  3504. patch.set_facecolor(color)
  3505. patch.set_alpha(0.75)
  3506. patch.set_edgecolor('black')
  3507. patch.set_linewidth(1.6)
  3508. ax.axhline(0, color='black', linewidth=0.6, linestyle='-', alpha=0.35)
  3509. # Staggered labels for all 8 bottom-panel significant genes
  3510. bottom_labels = [
  3511. ('GRX3', 0, 1.481, -0.55, 0.02), ('CYS4', 0, 1.247, 0.55, 0.02),
  3512. ('ADK1', 0, 1.159, -0.55, -0.02), ('PDX3', 0, 1.034, 0.55, -0.02),
  3513. ('ATP1', 2, 0.808, -0.55, 0.02), ('ERG4', 2, 0.510, 0.55, 0.02),
  3514. ('PDC6', 2, 0.459, -0.55, -0.02), ('COT1', 2, -0.510, 0.55, -0.02),
  3515. ]
  3516. for nm, gi, val, dx, dy in bottom_labels:
  3517. ax.annotate(nm, xy=(gi, val), xytext=(gi + dx, val + dy),
  3518. fontsize=9, fontweight='bold',
  3519. ha='right' if dx < 0 else 'left', va='center',
  3520. arrowprops=dict(arrowstyle='-', color='#333', lw=0.7,
  3521. shrinkA=2, shrinkB=2, alpha=0.6),
  3522. bbox=dict(boxstyle='round,pad=0.2', facecolor='white',
  3523. edgecolor='#999', linewidth=0.5, alpha=0.9), zorder=10)
  3524. ax.set_ylim(-0.75, BREAK + 0.1)
  3525. ax.set_xlim(-1.2, 4.5)
  3526. ax.set_xticks([0, 1, 2, 3])
  3527. ax.set_xticklabels([f'{l}\n(n={n[i]})' for i, l in enumerate(labels)],
  3528. fontsize=11, fontweight='bold')
  3529. ax.spines['top'].set_visible(False)
  3530. ax.spines['right'].set_visible(False)
  3531. ax.tick_params(labelsize=9)
  3532. ax.yaxis.set_major_locator(plt.MultipleLocator(0.25))
  3533. # ── BREAK SYMBOLS ────────────────────────────────────────────────────────
  3534. for ax_i in [ax_top, ax]:
  3535. ax_i.plot([-0.03, 0.03], [-0.02, 0.02], transform=ax_i.transAxes,
  3536. color='black', linewidth=1.2, clip_on=False)
  3537. ax_i.plot([-0.03, 0.03], [-0.10, -0.06], transform=ax_i.transAxes,
  3538. color='black', linewidth=1.2, clip_on=False)
  3539. ax.set_ylabel('Max signed thermal-protection score', fontsize=11, fontweight='bold')
  3540. fig.suptitle('EvoAge verdict vs wet-lab effect size — all significant genes labelled',
  3541. fontsize=13, fontweight='bold', y=0.995)
  3542. fig.text(0.5, 0.012,
  3543. 'Y-axis: thermal-protection score (KO - control difference) from the condition '
  3544. '(37 °C vs 30 °C OR Pulser vs 30 °C) with larger absolute value, sign kept.\n'
  3545. 'Positive = knockout improves survival (protective); negative = knockout reduces '
  3546. 'survival (sensitive). Axis break at 1.6. All 10 wet-lab significant genes labelled.',
  3547. ha='center', fontsize=8.5, style='italic', color='#444')
  3548. # Legend
  3549. leg = [Line2D([0], [0], marker='s', color='w', markerfacecolor=c,
  3550. markeredgecolor='black' if c not in ['#fff9c4', '#66bb6a'] else '#555',
  3551. markersize=12, label=l)
  3552. for l, c in [('Support', '#1b5e20'), ('Partial support', '#66bb6a'),
  3553. ('Weak support', '#fff9c4'), ('No support', '#d32f2f')]]
  3554. ax.legend(handles=leg, fontsize=9, loc='upper right', title='EvoAge verdict',
  3555. title_fontsize=10, bbox_to_anchor=(1.0, 1.06))
  3556. plt.tight_layout()
  3557. # ── Show in notebook ──────────────────────────────────────────────────
  3558. plt.show()
  3559. # ── Save figures ──────────────────────────────────────────────────────
  3560. for ext in ['svg', 'png', 'pdf']:
  3561. output_path = os.path.join(OUTPUT_DIR, f'boxplot_verdict_vs_effectsize_brokenaxis.{ext}')
  3562. plt.savefig(output_path, dpi=DPI, bbox_inches='tight')
  3563. print(f"✅ Saved: {output_path}")
  3564. plt.close()
  3565. # ── Print statistics ──────────────────────────────────────────────────
  3566. print("\n📊 Summary statistics:")
  3567. for i, (group, d) in enumerate(zip(labels, data)):
  3568. d_below = d[d <= BREAK]
  3569. d_above = d[d > BREAK]
  3570. print(f"\n{group.strip()}:")
  3571. print(f" Total: {len(d)} genes")
  3572. print(f" Below break (≤{BREAK}): {len(d_below)} genes")
  3573. print(f" Above break (>{BREAK}): {len(d_above)} genes")
  3574. if len(d) > 0:
  3575. print(f" Median: {np.median(d):.3f}")
  3576. print(f" Mean: {np.mean(d):.3f}")
  3577. print(f"\n✅ All done! Plots saved in: {OUTPUT_DIR}")
  3578. # %% [markdown]
  3579. # ## 5g
  3580. # %%
  3581. # This magic command makes plots appear in the notebook
  3582. %matplotlib inline
  3583. """
  3584. 04_bubble_table
  3585. ---------------
  3586. Integrated evidence matrix (bubble table) for the ten FDR-significant genes:
  3587. - Effect sizes (circle colour, diverging blue-red, -0.5 to +2.5)
  3588. - Statistical evidence (circle size, -log10 q, capped at 40 for display)
  3589. - EvoAge verdict pills
  3590. - Known / Novel / Reject / Dissent evidence counts
  3591. INPUT : /storage/Arushi/090526_EvoAge/kg_formation/all_figures/yeast-exp/EvoAge_plots/04_bubble_table/input_key_genes_evidence.csv
  3592. OUTPUT: EvoAge_Significant_Gene_Bubble_Table.svg / .png / .pdf in the same directory
  3593. """
  3594. import warnings
  3595. warnings.filterwarnings('ignore')
  3596. import pandas as pd
  3597. import numpy as np
  3598. import os
  3599. import matplotlib.pyplot as plt
  3600. from matplotlib.patches import Circle, Rectangle, FancyBboxPatch
  3601. from matplotlib.lines import Line2D
  3602. from matplotlib.colors import LinearSegmentedColormap
  3603. # ── Configuration ──────────────────────────────────────────────────────
  3604. # Input file path
  3605. INPUT_FILE = '/storage/Arushi/090526_EvoAge/kg_formation/all_figures/yeast-exp/EvoAge_plots/04_bubble_table/input_key_genes_evidence.csv'
  3606. # Output directory (same as input file location)
  3607. OUTPUT_DIR = os.path.dirname(INPUT_FILE)
  3608. DPI = 600
  3609. # ── Load data ──────────────────────────────────────────────────────────
  3610. print(f"Loading data from: {INPUT_FILE}")
  3611. df = pd.read_csv(INPUT_FILE)
  3612. print(f"Loaded {len(df)} key genes")
  3613. # Check required columns
  3614. required_cols = ['Name', 'Score_37C', 'Score_Pulser', 'neglog10q_37C', 'neglog10q_Pulser',
  3615. 'verdict', 'known', 'novel', 'reject', 'dissent']
  3616. missing_cols = [col for col in required_cols if col not in df.columns]
  3617. if missing_cols:
  3618. print(f"Warning: Missing columns: {missing_cols}")
  3619. print("Available columns:", df.columns.tolist())
  3620. # ── Prepare data ──────────────────────────────────────────────────────
  3621. # Order: both / 37C / pulser groups
  3622. order = ['CAX4', 'GRX3', 'ADK1', 'CYS4', 'DTD1', 'ATP1', 'COT1', 'ERG4', 'PDC6', 'PDX3']
  3623. group_of = {'CAX4': 'Both', 'GRX3': 'Both', 'ADK1': 'Both', 'CYS4': 'Both',
  3624. 'DTD1': '37 °C', 'ATP1': '37 °C', 'COT1': '37 °C', 'ERG4': '37 °C',
  3625. 'PDC6': '37 °C', 'PDX3': 'Pulser'}
  3626. # Ensure all genes are in the dataframe
  3627. for gene in order:
  3628. if gene not in df['Name'].values:
  3629. print(f"Warning: Gene {gene} not found in data")
  3630. df['order'] = df['Name'].map({n: i for i, n in enumerate(order)})
  3631. df = df.sort_values('order')
  3632. verdict_colors = {'support': '#5E35B1', 'weak_support': '#00897B', # dark purple, teal
  3633. 'partial_support': '#C0CA33', 'no_support': '#BDBDBD'} # muted yellow, grey
  3634. # ── Create figure ──────────────────────────────────────────────────────
  3635. fig, ax = plt.subplots(figsize=(13, 6.5))
  3636. ax.set_xlim(0, 13)
  3637. ax.set_ylim(-1, len(df) + 1.5)
  3638. # Column centres
  3639. col_score37, col_scoreP = 2.2, 3.6
  3640. col_q37, col_qP = 5.2, 6.6
  3641. col_verdict = 8.1
  3642. col_known, col_novel, col_reject, col_dissent = 9.4, 10.2, 11.0, 11.8
  3643. # Headers
  3644. for x, label in [(col_score37, '37 °C'), (col_scoreP, 'Pulser'),
  3645. (col_q37, '37 °C'), (col_qP, 'Pulser'),
  3646. (col_verdict, 'EvoAge verdict'),
  3647. (col_known, 'Known'), (col_novel, 'Novel'),
  3648. (col_reject, 'Reject'), (col_dissent, 'Dissent')]:
  3649. ax.text(x, len(df) + 0.7, label, ha='center', fontsize=9, fontweight='bold')
  3650. ax.text(2.9, len(df) + 1.3, 'Score', ha='center', fontsize=9, fontweight='bold')
  3651. ax.text(5.9, len(df) + 1.3, 'Statistical evidence', ha='center', fontsize=9, fontweight='bold')
  3652. # Group labels
  3653. group_y = {'Both': 4, '37 °C': 1.5, 'Pulser': -0.2}
  3654. for grp, y in group_y.items():
  3655. ax.text(-0.2, y, grp, ha='right', va='center', fontsize=10, fontweight='bold')
  3656. # Diverging colormap for effect size
  3657. cmap = LinearSegmentedColormap.from_list('div', ['#2166AC', '#F7F7F7', '#B2182B'])
  3658. norm_eff = plt.Normalize(-0.5, 2.5)
  3659. # ── Plot data ──────────────────────────────────────────────────────────
  3660. for i, (_, r) in enumerate(df.iterrows()):
  3661. y = len(df) - 1 - i # top-down
  3662. nm = r['Name']
  3663. ax.text(0.4, y, nm, ha='right', va='center', fontsize=9, fontweight='bold')
  3664. # Effect scores
  3665. for x, score in [(col_score37, r['Score_37C']), (col_scoreP, r['Score_Pulser'])]:
  3666. if pd.notna(score):
  3667. c = cmap(norm_eff(score))
  3668. ax.add_patch(Circle((x, y), 0.28, facecolor=c, edgecolor='black', lw=0.8, zorder=3))
  3669. ax.text(x, y, f'{score:+.2f}', ha='center', va='center', fontsize=7, fontweight='bold')
  3670. # Statistical evidence (-log10 q, capped at 40 for display)
  3671. for x, q in [(col_q37, r['neglog10q_37C']), (col_qP, r['neglog10q_Pulser'])]:
  3672. if pd.notna(q):
  3673. size = min(q, 40)
  3674. radius = 0.1 + 0.25 * (size / 40)
  3675. ax.add_patch(Circle((x, y), radius, facecolor='#7E57C2', alpha=0.5,
  3676. edgecolor='#4527A0', lw=0.8, zorder=3))
  3677. ax.text(x, y, f'{q:.1f}', ha='center', va='center', fontsize=7, fontweight='bold')
  3678. # Verdict pill
  3679. v = r['verdict']
  3680. vc = verdict_colors.get(v, '#BDBDBD')
  3681. ax.add_patch(FancyBboxPatch((col_verdict - 0.55, y - 0.22), 1.1, 0.44,
  3682. boxstyle='round,pad=0.02', facecolor=vc,
  3683. edgecolor='black', lw=0.6, zorder=3))
  3684. text_color = 'white' if v == 'support' else 'black'
  3685. ax.text(col_verdict, y, v.replace('_', ' '), ha='center', va='center',
  3686. fontsize=6.5, fontweight='bold', color=text_color)
  3687. # Evidence counts
  3688. for x, val, color in [(col_known, r['known'], '#C8E6C9'),
  3689. (col_novel, r['novel'], '#E1BEE7'),
  3690. (col_reject, r['reject'], '#B3E5FC'),
  3691. (col_dissent, r['dissent'], '#F8BBD0')]:
  3692. ax.add_patch(Rectangle((x - 0.32, y - 0.26), 0.64, 0.52, facecolor=color,
  3693. edgecolor='black', lw=0.5, zorder=3))
  3694. txt = str(int(val)) if pd.notna(val) else 'NA'
  3695. ax.text(x, y, txt, ha='center', va='center', fontsize=8, fontweight='bold')
  3696. ax.axis('off')
  3697. ax.set_title('Thermal-protection score | Triple support count',
  3698. fontsize=11, fontweight='bold', pad=10)
  3699. # Legend for q-size
  3700. for q in [5, 20, 40]:
  3701. ax.add_patch(Circle((11.0, -0.9), 0.1 + 0.25 * (q / 40), facecolor='#7E57C2',
  3702. alpha=0.5, edgecolor='#4527A0', lw=0.6))
  3703. ax.text(12.1, -0.9, '−log10(BH-FDR q), 40-max cap', fontsize=7, va='center')
  3704. # Add effect size colorbar legend
  3705. cbar_ax = fig.add_axes([0.92, 0.15, 0.02, 0.3])
  3706. cbar = plt.colorbar(plt.cm.ScalarMappable(norm=norm_eff, cmap=cmap),
  3707. cax=cbar_ax, orientation='vertical')
  3708. cbar.set_label('Effect size', fontsize=8)
  3709. cbar.ax.tick_params(labelsize=7)
  3710. plt.tight_layout()
  3711. # ── Show in notebook ──────────────────────────────────────────────────
  3712. plt.show()
  3713. # ── Save figures ──────────────────────────────────────────────────────
  3714. output_stem = os.path.join(OUTPUT_DIR, 'EvoAge_Significant_Gene_Bubble_Table')
  3715. for ext in ['svg', 'png', 'pdf']:
  3716. output_path = f"{output_stem}.{ext}"
  3717. plt.savefig(output_path, dpi=DPI, bbox_inches='tight')
  3718. print(f"✅ Saved: {output_path}")
  3719. plt.close()
  3720. # ── Print summary ──────────────────────────────────────────────────────
  3721. print("\n📊 Bubble Table Summary:")
  3722. print(f"Total key genes: {len(df)}")
  3723. print("\nGenes by group:")
  3724. for grp in ['Both', '37 °C', 'Pulser']:
  3725. count = sum([1 for gene in df['Name'] if group_of.get(gene) == grp])
  3726. print(f" {grp}: {count} genes")
  3727. print("\nVerdict distribution:")
  3728. for verdict in ['support', 'weak_support', 'partial_support', 'no_support']:
  3729. count = sum(df['verdict'] == verdict)
  3730. if count > 0:
  3731. print(f" {verdict}: {count} genes")
  3732. print(f"\n✅ All done! Plots saved in: {OUTPUT_DIR}")
  3733. # %% [markdown]
  3734. # ## Supplementary Figure 1
  3735. # %% [markdown]
  3736. # ## Supp 1a
  3737. # %%
  3738. # Make sure inline plotting is enabled
  3739. %matplotlib inline
  3740. import pandas as pd
  3741. import numpy as np
  3742. import matplotlib.pyplot as plt
  3743. import matplotlib.patches as mpatches
  3744. # Optional: If you don't need circlify anymore, remove the import
  3745. # import circlify
  3746. # -----------------------------------------------------------------------------
  3747. # 1. LOAD DATA
  3748. # -----------------------------------------------------------------------------
  3749. df = pd.read_csv("/storage/Arushi/090526_EvoAge/kg_formation/final_kg_building_3/STATS/NodeType_AllKGs_1to1_and_121_12M.csv")
  3750. # -----------------------------------------------------------------------------
  3751. # 2. PRETTY SHORT LABELS + COLOR PALETTE
  3752. # -----------------------------------------------------------------------------
  3753. short_names = {
  3754. "PMID": "PMID",
  3755. "Mutation": "mutation",
  3756. "ChemicalEntity": "chemical",
  3757. "Protein": "protein",
  3758. "Phenotype": "phenotype",
  3759. "Gene": "gene",
  3760. "Disease": "disease",
  3761. "BiologicalProcess": "biological",
  3762. "AnatomicalEntity": "anatomy",
  3763. "MolecularFunction": "molecular",
  3764. "PlantSpecies": "plant",
  3765. "Tissue": "tissue",
  3766. "Pathway": "pathway",
  3767. "CellularComponent": "cellular",
  3768. "Mirna": "mirna",
  3769. "Species": "species",
  3770. }
  3771. nodetype_colors = {
  3772. "PMID": "#c8d89a",
  3773. "Mutation": "#b0b0b0",
  3774. "ChemicalEntity": "#c9a86a",
  3775. "Protein": "#dcb0c4",
  3776. "Phenotype": "#e8a4b5",
  3777. "Gene": "#6e85b0",
  3778. "Disease": "#a89bc8",
  3779. "BiologicalProcess": "#8ec4de",
  3780. "AnatomicalEntity": "#e8a586",
  3781. "MolecularFunction": "#c9a86a",
  3782. "PlantSpecies": "#a67aa8",
  3783. "Tissue": "#f0a878",
  3784. "Pathway": "#b39ddb",
  3785. "CellularComponent": "#f5d896",
  3786. "Mirna": "#b06ba0",
  3787. "Species": "#dcdc82",
  3788. }
  3789. # -----------------------------------------------------------------------------
  3790. # 3. CREATE THE PLOT
  3791. # -----------------------------------------------------------------------------
  3792. # Clear any existing figures
  3793. plt.close('all')
  3794. # Create figure and axis
  3795. fig, ax1 = plt.subplots(1, 1, figsize=(12, 8))
  3796. # =============================================================================
  3797. # BAR PLOT — EvoAge_121_12M
  3798. # =============================================================================
  3799. evo = df[["NodeType", "EvoAge_121_12M"]].copy()
  3800. evo = evo[evo["EvoAge_121_12M"] > 0].sort_values("EvoAge_121_12M", ascending=False).reset_index(drop=True)
  3801. evo["log"] = np.log10(evo["EvoAge_121_12M"] + 1)
  3802. x = np.arange(len(evo))
  3803. bar_colors = [nodetype_colors.get(nt, "#cccccc") for nt in evo["NodeType"]]
  3804. # Create the bars
  3805. bars = ax1.bar(x, evo["log"], color=bar_colors, edgecolor="black", linewidth=0.4, width=0.75)
  3806. # Count labels rotated 90° inside/on each bar
  3807. for xi, (_, row) in zip(x, evo.iterrows()):
  3808. ax1.text(xi, row["log"] * 0.5, f"{int(row['EvoAge_121_12M'])}",
  3809. ha="center", va="center", rotation=90, fontsize=9, color="black")
  3810. # Set labels and title
  3811. ax1.set_xticks(x)
  3812. ax1.set_xticklabels([short_names.get(nt, nt) for nt in evo["NodeType"]],
  3813. rotation=90, fontsize=10)
  3814. ax1.set_ylabel(r"$\log_{10}$(count)", fontsize=13)
  3815. ax1.set_title("Node count (EvoAge)", fontsize=15, pad=10)
  3816. # Remove top and right spines
  3817. ax1.spines["top"].set_visible(False)
  3818. ax1.spines["right"].set_visible(False)
  3819. ax1.tick_params(axis="both", which="major", length=5, width=0.8, color="black", labelsize=10)
  3820. ax1.grid(False)
  3821. # Adjust layout
  3822. plt.tight_layout()
  3823. # -----------------------------------------------------------------------------
  3824. # 4. DISPLAY THE PLOT
  3825. # -----------------------------------------------------------------------------
  3826. # Show the plot
  3827. plt.show()
  3828. # -----------------------------------------------------------------------------
  3829. # 5. SAVE AS PDF (after showing)
  3830. # -----------------------------------------------------------------------------
  3831. # Save the figure
  3832. fig.savefig("FIG1/Fig1_supp_evoage_bar.pdf", format="pdf", bbox_inches="tight")
  3833. print("Plot saved successfully!")
  3834. # %% [markdown]
  3835. # ## Supp 1b
  3836. # %%
  3837. import numpy as np
  3838. import pandas as pd
  3839. import matplotlib.pyplot as plt
  3840. from matplotlib.path import Path
  3841. from matplotlib.patches import PathPatch, Wedge, Rectangle
  3842. import matplotlib as mpl
  3843. # ------------------------------------------------------------------
  3844. # 0. Editable SVG text
  3845. # ------------------------------------------------------------------
  3846. mpl.rcParams['svg.fonttype'] = 'none'
  3847. mpl.rcParams['font.family'] = 'DejaVu Sans'
  3848. # ------------------------------------------------------------------
  3849. # 1. Load & parse
  3850. # ------------------------------------------------------------------
  3851. CSV_PATH = 'RelationType_AllKGs_1to1_and_121_12M.csv'
  3852. VALUE_COL = 'EvoAge_121_12M'
  3853. df = pd.read_csv(CSV_PATH)
  3854. df = df[df['Relation'] != 'Total Triples'].copy()
  3855. df = df[df[VALUE_COL] > 0].reset_index(drop=True)
  3856. def parse_relation(relation):
  3857. parts = relation.split('_')
  3858. if len(parts) == 2:
  3859. return parts[0], None, parts[1], 'entity'
  3860. else:
  3861. # first token = source, LAST token = target, everything in between = verb
  3862. return parts[0], '_'.join(parts[1:-1]), parts[-1], 'verb'
  3863. parsed = df['Relation'].apply(parse_relation)
  3864. df['src'] = parsed.apply(lambda x: x[0])
  3865. df['verb'] = parsed.apply(lambda x: x[1])
  3866. df['tgt'] = parsed.apply(lambda x: x[2])
  3867. df['edge_type'] = parsed.apply(lambda x: x[3])
  3868. # camel-case -> "spaced lower case" for nice legend text, e.g.
  3869. # "NegativelyAssociatedWith" -> "negatively associated with"
  3870. def humanize(camel):
  3871. out = []
  3872. for ch in camel:
  3873. if ch.isupper() and out:
  3874. out.append(' ')
  3875. out.append(ch.lower())
  3876. return ''.join(out)
  3877. # ------------------------------------------------------------------
  3878. # 2. log10 weight used for ribbon/arc geometry (raw count kept for labels)
  3879. # ------------------------------------------------------------------
  3880. df['raw'] = df[VALUE_COL]
  3881. df['logw'] = np.log10(df['raw'] + 1)
  3882. # ------------------------------------------------------------------
  3883. # 3. Entities & node ordering
  3884. # ------------------------------------------------------------------
  3885. entities = sorted(set(df['src']) | set(df['tgt']))
  3886. idx = {e: i for i, e in enumerate(entities)}
  3887. n = len(entities)
  3888. # total log-weighted flux touching each entity (for arc size)
  3889. flux = np.zeros(n)
  3890. for _, row in df.iterrows():
  3891. flux[idx[row['src']]] += row['logw']
  3892. flux[idx[row['tgt']]] += row['logw']
  3893. order = np.argsort(-flux)
  3894. entities = [entities[i] for i in order]
  3895. idx = {e: i for i, e in enumerate(entities)}
  3896. flux = flux[order]
  3897. # ------------------------------------------------------------------
  3898. # 4. Colour palettes
  3899. # - entity palette: pastel, one colour per entity (used for entity-entity
  3900. # ribbons AND node wedges)
  3901. # - relation-type palette: separate, more saturated colours, used ONLY for
  3902. # verb-qualified ribbons (Promotes / Inhibits / etc.)
  3903. # ------------------------------------------------------------------
  3904. # Colours matched (pixel-sampled) from the reference "EvoAge (Node count)"
  3905. # bar chart, so the entity colours are consistent across figures.
  3906. entity_hex = {
  3907. 'PMID': '#C8D89A', # olive-green
  3908. 'Mutation': '#B0B0B0', # grey
  3909. 'ChemicalEntity': '#C9A86A', # tan / gold-brown
  3910. 'Protein': '#DCB0C4', # pink
  3911. 'Phenotype': '#E8A4B5', # rose pink
  3912. 'Gene': '#6E85B0', # dark blue / navy
  3913. 'AnatomicalEntity': '#E8A586', # salmon / orange
  3914. 'Disease': '#A89BC8', # light purple
  3915. 'BiologicalProcess': '#8EC4DE', # sky blue
  3916. 'MolecularFunction': '#B08968', # muted brown (reference re-used the
  3917. # chemical tan here; shifted slightly
  3918. # darker so every entity stays
  3919. # distinguishable, per your note)
  3920. 'PlantSpecies': '#A67AA8', # plum
  3921. 'Tissue': '#F0A878', # orange
  3922. 'Pathway': '#B39DDB', # lavender-purple
  3923. 'CellularComponent': '#F5D896', # cream / pale yellow
  3924. 'Mirna': '#B06BA0', # magenta-plum
  3925. 'Species': '#DCDC82', # yellow-green
  3926. 'Nodes': '#C0504D', # brick red (not in the reference
  3927. # chart, so a new but equally
  3928. # contrasting colour was added)
  3929. }
  3930. node_color = {e: entity_hex.get(e, '#CCCCCC') for e in entities}
  3931. # Relation-type (verb-qualified edge) palette: kept intentionally as
  3932. # saturated/dark "jewel tones" -- a different visual register from the
  3933. # pastel entity colours above -- so verb-qualified ribbons never get
  3934. # confused with a plain entity-entity ribbon, even where hue families
  3935. # are close (e.g. reds vs pinks, greens vs olive).
  3936. verb_palette = {
  3937. 'AssociatedWith': '#073B3A', # deep teal
  3938. 'NoEffect': '#6A00FF', # vivid violet
  3939. 'Promotes': '#003566', # dark navy
  3940. 'Inhibits': '#9D0208', # dark red
  3941. 'NegativelyAssociatedWith': '#2B9348', # deep green
  3942. 'PositivelyAssociatedWith': '#BC6C25', # burnt amber
  3943. 'NotAssociatedWith': '#4A4E69', # slate grey-purple
  3944. }
  3945. def edge_color(row):
  3946. if row['edge_type'] == 'entity':
  3947. return node_color[row['src']]
  3948. return verb_palette.get(row['verb'], '#999999')
  3949. df['color'] = df.apply(edge_color, axis=1)
  3950. # short display names for entity labels (match reference style, e.g.
  3951. # "BiologicalProcess" -> "Biological")
  3952. short_name = {
  3953. 'BiologicalProcess': 'Biological',
  3954. 'CellularComponent': 'Cellular',
  3955. 'MolecularFunction': 'Molecular',
  3956. 'AnatomicalEntity': 'Anatomical',
  3957. 'ChemicalEntity': 'Chemical',
  3958. 'PlantSpecies': 'Plant Species',
  3959. }
  3960. def short(e):
  3961. return short_name.get(e, e)
  3962. # ------------------------------------------------------------------
  3963. # 5. Layout: angles for each node's arc (deg), clockwise from top
  3964. # ------------------------------------------------------------------
  3965. GAP_DEG = 2.4
  3966. total_gap = GAP_DEG * n
  3967. available_deg = 360 - total_gap
  3968. frac = flux / flux.sum()
  3969. node_span = frac * available_deg
  3970. start_ang = 90.0
  3971. node_start, node_end = {}, {}
  3972. cursor = start_ang
  3973. for e, span in zip(entities, node_span):
  3974. node_start[e] = cursor - span
  3975. node_end[e] = cursor
  3976. cursor = cursor - span - GAP_DEG
  3977. node_center = {e: (node_start[e] + node_end[e]) / 2 for e in entities}
  3978. # ------------------------------------------------------------------
  3979. # 6. Slots (one per ribbon end) sized by log-weight
  3980. # ------------------------------------------------------------------
  3981. slots = {e: [] for e in entities}
  3982. for edge_id, row in df.iterrows():
  3983. s, t, w = row['src'], row['tgt'], row['logw']
  3984. slots[s].append(dict(partner=t, weight=w, edge_id=edge_id, end='src'))
  3985. slots[t].append(dict(partner=s, weight=w, edge_id=edge_id, end='tgt'))
  3986. for e in entities:
  3987. def sort_key(sl):
  3988. p = sl['partner']
  3989. return -1000 if p == e else node_center[p]
  3990. slots[e].sort(key=sort_key)
  3991. edge_geom = {}
  3992. for e in entities:
  3993. s_list = slots[e]
  3994. wsum = sum(sl['weight'] for sl in s_list)
  3995. if wsum == 0:
  3996. continue
  3997. span = node_end[e] - node_start[e]
  3998. cur = node_start[e]
  3999. for sl in s_list:
  4000. w = span * (sl['weight'] / wsum)
  4001. a0, a1 = cur, cur + w
  4002. cur = a1
  4003. edge_geom.setdefault(sl['edge_id'], {})[sl['end']] = (a0, a1)
  4004. # ------------------------------------------------------------------
  4005. # 7. Drawing helpers
  4006. # ------------------------------------------------------------------
  4007. R_IN = 1.0
  4008. R_OUT = 1.09
  4009. N_ARC_PTS = 40
  4010. def arc_points(a0, a1, r=R_IN):
  4011. angs = np.linspace(np.radians(a0), np.radians(a1), N_ARC_PTS)
  4012. return np.column_stack([r * np.cos(angs), r * np.sin(angs)])
  4013. def ribbon_path(srcA, srcB, tgtA, tgtB):
  4014. p_src = arc_points(srcA, srcB)
  4015. p_tgt = arc_points(tgtA, tgtB)
  4016. verts = [p_src[0]]
  4017. codes = [Path.MOVETO]
  4018. for p in p_src[1:]:
  4019. verts.append(p); codes.append(Path.LINETO)
  4020. c1, c2 = p_src[-1] * 0.35, p_tgt[0] * 0.35
  4021. verts += [c1, c2, p_tgt[0]]; codes += [Path.CURVE4]*3
  4022. for p in p_tgt[1:]:
  4023. verts.append(p); codes.append(Path.LINETO)
  4024. c3, c4 = p_tgt[-1] * 0.35, p_src[0] * 0.35
  4025. verts += [c3, c4, p_src[0]]; codes += [Path.CURVE4]*3
  4026. return Path(verts, codes)
  4027. # ------------------------------------------------------------------
  4028. # 8. Plot
  4029. # ------------------------------------------------------------------
  4030. fig, ax = plt.subplots(figsize=(14, 12), subplot_kw=dict(aspect='equal'))
  4031. ax.set_xlim(-1.85, 2.35)
  4032. ax.set_ylim(-1.55, 1.65)
  4033. ax.axis('off')
  4034. edges_sorted = df.sort_values('raw', ascending=False).index.tolist()
  4035. for eid in edges_sorted:
  4036. row = df.loc[eid]
  4037. if eid not in edge_geom or 'src' not in edge_geom[eid] or 'tgt' not in edge_geom[eid]:
  4038. continue
  4039. a0, a1 = edge_geom[eid]['src']
  4040. b0, b1 = edge_geom[eid]['tgt']
  4041. path = ribbon_path(a0, a1, b0, b1)
  4042. patch = PathPatch(path, facecolor=row['color'], edgecolor='none', alpha=0.65,
  4043. linewidth=0, zorder=1)
  4044. ax.add_patch(patch)
  4045. # node wedges
  4046. for e in entities:
  4047. w = Wedge((0, 0), R_OUT, node_start[e], node_end[e], width=R_OUT - R_IN,
  4048. facecolor=node_color[e], edgecolor='white', linewidth=1.2, zorder=3)
  4049. ax.add_patch(w)
  4050. # raw-count numbers next to each ribbon slot
  4051. for eid, ends in edge_geom.items():
  4052. raw = df.loc[eid, 'raw']
  4053. label = f'{int(raw):,}'
  4054. for end_key in ('src', 'tgt'):
  4055. if end_key not in ends:
  4056. continue
  4057. a0, a1 = ends[end_key]
  4058. mid = np.radians((a0 + a1) / 2)
  4059. r_txt = R_OUT + 0.025
  4060. x, y = r_txt * np.cos(mid), r_txt * np.sin(mid)
  4061. rot = np.degrees(mid)
  4062. ha = 'left'
  4063. if 90 < (np.degrees(mid) % 360) < 270:
  4064. rot += 180
  4065. ha = 'right'
  4066. ax.text(x, y, label, rotation=rot, rotation_mode='anchor',
  4067. ha=ha, va='center', fontsize=6.3, zorder=5)
  4068. # entity name labels (outer ring)
  4069. for e in entities:
  4070. ang = np.radians(node_center[e])
  4071. r_label = R_OUT + 0.16
  4072. x, y = r_label * np.cos(ang), r_label * np.sin(ang)
  4073. rot = node_center[e]
  4074. ha = 'left'
  4075. if 90 < node_center[e] % 360 < 270:
  4076. rot += 180
  4077. ha = 'right'
  4078. ax.text(x, y, short(e), rotation=rot, rotation_mode='anchor',
  4079. ha=ha, va='center', fontsize=11, fontweight='medium', zorder=4)
  4080. # "log10(count)" radial axis label (reference-style)
  4081. ax.text(-1.42, 0, r'$\log_{10}(\mathrm{count})$', rotation=90, rotation_mode='anchor',
  4082. ha='center', va='center', fontsize=13)
  4083. ax.set_title('EvoAge Knowledge Graph — Entity Relation Chord Diagram\n'
  4084. 'relation counts (EvoAge, 121:12M) | ribbon width \u221d log$_{10}$(count)',
  4085. fontsize=13, pad=10)
  4086. # ------------------------------------------------------------------
  4087. # 9. Legend: entities block + relation-type block
  4088. # ------------------------------------------------------------------
  4089. leg_x = 1.42
  4090. leg_y = 1.55
  4091. box_w, box_h, gap = 0.075, 0.055, 0.018
  4092. ax.text(leg_x, leg_y + 0.06, 'Entities', fontsize=12.5, fontweight='bold', va='center')
  4093. y = leg_y
  4094. for e in entities:
  4095. ax.add_patch(Rectangle((leg_x, y - box_h/2), box_w, box_h,
  4096. facecolor=node_color[e], edgecolor='none'))
  4097. ax.text(leg_x + box_w + 0.03, y, short(e), fontsize=10.5, va='center')
  4098. y -= (box_h + gap)
  4099. y -= 0.04
  4100. ax.text(leg_x, y + 0.06, 'Relation type\n(qualified edges, separate colour scheme)',
  4101. fontsize=11, fontweight='bold', va='center')
  4102. y -= 0.02
  4103. for verb, col in verb_palette.items():
  4104. if verb not in df.loc[df['edge_type']=='verb', 'verb'].unique():
  4105. continue
  4106. ax.add_patch(Rectangle((leg_x, y - box_h/2), box_w, box_h,
  4107. facecolor=col, edgecolor='none'))
  4108. ax.text(leg_x + box_w + 0.03, y, humanize(verb), fontsize=10.5, va='center')
  4109. y -= (box_h + gap)
  4110. plt.tight_layout()
  4111. plt.savefig('EvoAge_ChordDiagram_log10_matched.svg', format='svg', bbox_inches='tight')
  4112. plt.savefig('EvoAge_ChordDiagram_log10_matched_preview.png', format='png', dpi=170, bbox_inches='tight')
  4113. print("Saved EvoAge_ChordDiagram_log10_matched.svg and preview png")
  4114. # %% [markdown]
  4115. # ## Supp 1c
  4116. # %%
  4117. import pandas as pd
  4118. import plotly.graph_objects as go
  4119. # Load data
  4120. df = pd.read_csv('orthology-mapping.csv')
  4121. print("Data loaded successfully!")
  4122. print(df)
  4123. # Prepare data
  4124. species_list = df['Species'].tolist()
  4125. labels = species_list + ['1-to-1\northolog', '1-to-M\northologs', 'Unmapped']
  4126. source = []
  4127. target = []
  4128. value = []
  4129. species_indices = {species: idx for idx, species in enumerate(species_list)}
  4130. category_indices = {
  4131. 'OneToOne': len(species_list),
  4132. 'OneToMany': len(species_list) + 1,
  4133. 'Unmapped': len(species_list) + 2
  4134. }
  4135. for idx, row in df.iterrows():
  4136. species = row['Species']
  4137. species_idx = species_indices[species]
  4138. source.append(species_idx)
  4139. target.append(category_indices['OneToOne'])
  4140. value.append(row['OneToOne'])
  4141. source.append(species_idx)
  4142. target.append(category_indices['OneToMany'])
  4143. value.append(row['OneToMany'])
  4144. source.append(species_idx)
  4145. target.append(category_indices['Unmapped'])
  4146. value.append(row['Unmapped'])
  4147. total_1to1 = df['OneToOne'].sum()
  4148. total_1toM = df['OneToMany'].sum()
  4149. total_unmapped = df['Unmapped'].sum()
  4150. # Cleaner version with gray links
  4151. fig = go.Figure(data=[go.Sankey(
  4152. node=dict(
  4153. pad=30,
  4154. thickness=25,
  4155. line=dict(color="white", width=1),
  4156. label=labels,
  4157. color=['#3498DB', '#E67E22', '#2ECC71', '#E74C3C', '#9B59B6',
  4158. '#2980B9', '#F39C12', '#C0392B'],
  4159. x=[0.05, 0.05, 0.05, 0.05, 0.05, 0.82, 0.82, 0.82],
  4160. y=[0.15, 0.32, 0.49, 0.66, 0.83, 0.25, 0.55, 0.80]
  4161. ),
  4162. link=dict(
  4163. source=source,
  4164. target=target,
  4165. value=value,
  4166. color='rgba(180, 180, 180, 0.5)' # Gray semi-transparent links
  4167. )
  4168. )])
  4169. fig.update_layout(
  4170. title=dict(
  4171. text="<b>Ortholog Gene Mapping from Other Species to Human</b>",
  4172. font=dict(size=22, family="Arial", color="#2C3E50"),
  4173. x=0.5
  4174. ),
  4175. width=1200,
  4176. height=700,
  4177. margin=dict(t=80, b=120, l=60, r=60),
  4178. plot_bgcolor='white',
  4179. paper_bgcolor='white'
  4180. )
  4181. # Add annotations
  4182. fig.add_annotation(
  4183. x=0.5,
  4184. y=-0.10,
  4185. xref="paper",
  4186. yref="paper",
  4187. text=(f"<b>Summary</b> | "
  4188. f"<span style='color:#2980B9;'>●</span> 1-to-1: <b>{total_1to1:,}</b> | "
  4189. f"<span style='color:#F39C12;'>●</span> 1-to-M: <b>{total_1toM:,}</b> | "
  4190. f"<span style='color:#C0392B;'>●</span> Unmapped: <b>{total_unmapped:,}</b>"),
  4191. showarrow=False,
  4192. font=dict(size=14),
  4193. align="center"
  4194. )
  4195. fig.show()
  4196. fig.write_html("sankey_ortholog_mapping_simple.html")
  4197. print("\n✅ Simple Sankey diagram saved as 'sankey_ortholog_mapping_simple.html'")
  4198. # %% [markdown]
  4199. # ## Supp 1e
  4200. # %%
  4201. import pandas as pd
  4202. import matplotlib.pyplot as plt
  4203. import numpy as np
  4204. from matplotlib.patches import Rectangle
  4205. import matplotlib.colors as mcolors
  4206. # Read the CSV file
  4207. df = pd.read_csv('ALL_data_Source_Triple_source_databases_KG.csv')
  4208. # Filter for aging data sources only (KG_Type == 'Aging')
  4209. aging_df = df[df['KG_Type'] == 'Aging'].copy()
  4210. # Define all columns to check (all relationship types)
  4211. all_columns = aging_df.columns.tolist()
  4212. # Remove non-data columns
  4213. exclude_cols = ['Source', 'KG_Type', 'Total', 'Species_AssociatedWith', 'PlantSpecies_ChemicalEntity',
  4214. 'PlantSpecies_Disease', 'PMID_CellularComponent', 'PMID_ChemicalEntity', 'PMID_Disease',
  4215. 'PMID_Protein', 'PMID_Tissue', 'Phenotype_BiologicalProcess', 'Phenotype_CellularComponent',
  4216. 'Phenotype_MolecularFunction', 'Phenotype_Protein', 'Phenotype_Phenotype', 'Phenotype_Gene',
  4217. 'Phenotype_Disease', 'Phenotype_ChemicalEntity', 'Mirna_Gene']
  4218. columns_to_plot = [col for col in all_columns if col not in exclude_cols and col != 'Total']
  4219. # Filter to columns that have some data in aging sources
  4220. valid_columns = []
  4221. for col in columns_to_plot:
  4222. if col in aging_df.columns and aging_df[col].sum() > 0:
  4223. valid_columns.append(col)
  4224. # Prepare the data matrix
  4225. aging_data = aging_df[['Source'] + valid_columns].copy()
  4226. aging_data.set_index('Source', inplace=True)
  4227. # Remove sources that have no data in these columns
  4228. aging_data = aging_data[aging_data.sum(axis=1) > 0]
  4229. # Replace 0 with NaN (we don't want to display zeros)
  4230. aging_data = aging_data.replace(0, np.nan)
  4231. # Apply log10 transformation (adding small constant to avoid log(0) issues)
  4232. # We'll use log10(count + 1) for the color mapping
  4233. log_data = np.log10(aging_data.values + 1)
  4234. # The +1 ensures log(1)=0 for zero values, but we'll mask NaNs anyway
  4235. # Create the plot with the specific color gradient
  4236. fig, ax = plt.subplots(figsize=(25, 8))
  4237. # Define the colormap matching the screenshot (blue to red)
  4238. # This creates a diverging colormap from blue through white to red
  4239. cmap = plt.cm.viridis
  4240. # Create the heatmap
  4241. # We use log_data for color mapping but keep original values for text
  4242. im = ax.imshow(log_data, cmap=cmap, aspect='auto',
  4243. vmin=np.nanmin(log_data), vmax=np.nanmax(log_data))
  4244. # Set ticks and labels
  4245. ax.set_xticks(np.arange(len(aging_data.columns)))
  4246. ax.set_yticks(np.arange(len(aging_data.index)))
  4247. ax.set_xticklabels(aging_data.columns, rotation=90, ha='center', fontsize=9)
  4248. ax.set_yticklabels(aging_data.index, fontsize=10)
  4249. # Add colored boxes with borders for each cell
  4250. for i in range(len(aging_data.index)):
  4251. for j in range(len(aging_data.columns)):
  4252. value = aging_data.iloc[i, j]
  4253. if not pd.isna(value):
  4254. # Add rectangle border
  4255. rect = Rectangle((j-0.5, i-0.5), 1, 1, fill=False, edgecolor='black', linewidth=1.5)
  4256. ax.add_patch(rect)
  4257. # Format the value
  4258. if value >= 1000000:
  4259. text = f'{value/1000000:.1f}M'
  4260. elif value >= 1000:
  4261. text = f'{value/1000:.1f}K'
  4262. else:
  4263. text = f'{int(value)}'
  4264. # Choose text color based on background
  4265. log_val = log_data[i, j]
  4266. if log_val > np.nanmax(log_data) * 0.4:
  4267. color = 'white'
  4268. else:
  4269. color = 'black'
  4270. ax.text(j, i, text, ha='center', va='center', color=color, fontsize=7, weight='bold')
  4271. # Add colorbar with log scale labels
  4272. cbar = plt.colorbar(im, ax=ax, shrink=0.8)
  4273. cbar.set_label('log₁₀(count + 1)', fontsize=12)
  4274. # Add legend for log values
  4275. log_ticks = np.arange(0, np.nanmax(log_data) + 0.5, 0.5)
  4276. cbar.set_ticks(log_ticks)
  4277. cbar.set_ticklabels([f'{tick:.1f}' for tick in log_ticks])
  4278. # Set title
  4279. ax.set_title('Aging Data Sources - Relationship Types (Log Scale)', fontsize=16, pad=20)
  4280. # Adjust layout
  4281. plt.tight_layout()
  4282. # Save the figure (optional)
  4283. # plt.savefig('aging_data_heatmap.png', dpi=300, bbox_inches='tight')
  4284. plt.show()
  4285. # Print summary statistics
  4286. print("Aging Data Sources Summary:")
  4287. print("=" * 60)
  4288. print(f"Number of aging data sources: {len(aging_data.index)}")
  4289. print(f"Data sources: {', '.join(aging_data.index)}")
  4290. print(f"\nNumber of relationship types with data: {len(aging_data.columns)}")
  4291. print("\nTop relationship types by total count:")
  4292. totals = aging_data.sum(axis=0).sort_values(ascending=False)
  4293. for col, val in totals.head(10).items():
  4294. print(f" {col}: {int(val):,}")
  4295. # Alternative: Create a version that only shows columns where any source has data
  4296. def create_filtered_heatmap():
  4297. """Create a more focused heatmap with only active relationship types"""
  4298. # Get columns where at least one source has data
  4299. active_columns = []
  4300. for col in valid_columns:
  4301. if aging_df[col].sum() > 0:
  4302. active_columns.append(col)
  4303. # Prepare data
  4304. filtered_data = aging_df[['Source'] + active_columns].copy()
  4305. filtered_data.set_index('Source', inplace=True)
  4306. filtered_data = filtered_data.replace(0, np.nan)
  4307. filtered_data = filtered_data[filtered_data.sum(axis=1) > 0]
  4308. # Log transform
  4309. log_filtered = np.log10(filtered_data.values + 1)
  4310. # Plot
  4311. fig, ax = plt.subplots(figsize=(20, 8))
  4312. im = ax.imshow(log_filtered, cmap='RdYlBu_r', aspect='auto',
  4313. vmin=0, vmax=np.nanmax(log_filtered))
  4314. ax.set_xticks(np.arange(len(filtered_data.columns)))
  4315. ax.set_yticks(np.arange(len(filtered_data.index)))
  4316. ax.set_xticklabels(filtered_data.columns, rotation=90, ha='center', fontsize=8)
  4317. ax.set_yticklabels(filtered_data.index, fontsize=10)
  4318. # Add boxes and text
  4319. for i in range(len(filtered_data.index)):
  4320. for j in range(len(filtered_data.columns)):
  4321. value = filtered_data.iloc[i, j]
  4322. if not pd.isna(value):
  4323. rect = Rectangle((j-0.5, i-0.5), 1, 1, fill=False, edgecolor='black', linewidth=1.5)
  4324. ax.add_patch(rect)
  4325. if value >= 1000000:
  4326. text = f'{value/1000000:.1f}M'
  4327. elif value >= 1000:
  4328. text = f'{value/1000:.1f}K'
  4329. else:
  4330. text = f'{int(value)}'
  4331. log_val = log_filtered[i, j]
  4332. color = 'white' if log_val > np.nanmax(log_filtered) * 0.4 else 'black'
  4333. ax.text(j, i, text, ha='center', va='center', color=color, fontsize=7, weight='bold')
  4334. cbar = plt.colorbar(im, ax=ax, shrink=0.8)
  4335. cbar.set_label('log₁₀(count + 1)', fontsize=12)
  4336. ax.set_title('Aging Data Sources - Active Relationship Types (Log Scale)', fontsize=14, pad=15)
  4337. plt.tight_layout()
  4338. plt.show()
  4339. # Uncomment to see the filtered version
  4340. # create_filtered_heatmap()
  4341. # Create a table version showing only the non-zero values
  4342. def create_table_view():
  4343. """Create a table view similar to the screenshot"""
  4344. # Get the data
  4345. table_data = aging_df[['Source'] + valid_columns].copy()
  4346. table_data.set_index('Source', inplace=True)
  4347. # Only keep columns with data
  4348. cols_with_data = table_data.columns[table_data.sum(axis=0) > 0]
  4349. table_data = table_data[cols_with_data]
  4350. # Only keep rows with data
  4351. rows_with_data = table_data.index[table_data.sum(axis=1) > 0]
  4352. table_data = table_data.loc[rows_with_data]
  4353. # Display as a styled table
  4354. fig, ax = plt.subplots(figsize=(20, len(table_data) * 0.5 + 1))
  4355. ax.axis('off')
  4356. # Create the table
  4357. table = ax.table(cellText=table_data.values.astype(str).tolist(),
  4358. rowLabels=table_data.index,
  4359. colLabels=table_data.columns,
  4360. cellLoc='center',
  4361. rowLoc='center',
  4362. loc='center')
  4363. table.auto_set_font_size(False)
  4364. table.set_fontsize(8)
  4365. table.scale(1.2, 1.5)
  4366. # Color the cells based on values
  4367. for i in range(len(table_data.index)):
  4368. for j in range(len(table_data.columns)):
  4369. value = table_data.iloc[i, j]
  4370. if value > 0:
  4371. # Get cell
  4372. cell = table[(i+1, j)]
  4373. # Color based on log value
  4374. log_val = np.log10(value + 1)
  4375. max_log = np.log10(table_data.max().max() + 1)
  4376. if log_val / max_log > 0.6:
  4377. cell.set_facecolor('#e74c3c') # Red
  4378. cell.set_text_props(color='white')
  4379. elif log_val / max_log > 0.3:
  4380. cell.set_facecolor('#f39c12') # Orange
  4381. cell.set_text_props(color='white')
  4382. else:
  4383. cell.set_facecolor('#3498db') # Blue
  4384. cell.set_text_props(color='white')
  4385. else:
  4386. cell = table[(i+1, j)]
  4387. cell.set_facecolor('#f0f0f0')
  4388. cell.set_text_props(color='gray')
  4389. plt.title('Aging Data Sources - Relationship Counts', fontsize=14, pad=20)
  4390. plt.tight_layout()
  4391. plt.show()
  4392. # Uncomment to see the table view
  4393. # create_table_view()
  4394. # %% [markdown]
  4395. # ## Supp 1f
  4396. # %%
  4397. import pandas as pd
  4398. import matplotlib.pyplot as plt
  4399. import numpy as np
  4400. from matplotlib.patches import Rectangle
  4401. import matplotlib.colors as mcolors
  4402. # Read the CSV file
  4403. df = pd.read_csv('Triple_source_databases_KG.csv')
  4404. # Filter for generalised data sources only (KG_Type == 'Generalised')
  4405. generalised_df = df[df['KG_Type'] == 'Generalised'].copy()
  4406. # Define the exact order of sources from the screenshot
  4407. source_order = [
  4408. 'BOCK',
  4409. 'BindingDB',
  4410. 'BioGrakn',
  4411. 'BioSNAP',
  4412. 'Biogrid', # Note: BioGRID in the list
  4413. 'CKG',
  4414. 'CROssBAR',
  4415. 'Chembl', # Note: ChEMBL in the list
  4416. 'DRKG',
  4417. 'DTINet',
  4418. 'DrugBank',
  4419. 'Evolf', # Note: EvOf in the list
  4420. 'FICE',
  4421. 'FlyBase',
  4422. 'Harmonizome',
  4423. 'Hetionet',
  4424. 'IMPPAT',
  4425. 'MGI_DO', # Note: MGI in the list
  4426. 'MonarchKG',
  4427. 'Mouse_Net', # Note: MouseNet in the list
  4428. 'PharmKG',
  4429. 'PheKnowLator',
  4430. 'Phytohub',
  4431. 'PrimeKG',
  4432. 'SGD',
  4433. 'SMS',
  4434. 'STITCH',
  4435. 'STRING',
  4436. 'TARKG',
  4437. 'TTD',
  4438. 'Worm_Interactome_Database', # Note: WIDB in the list
  4439. 'YeastNet',
  4440. 'ZFIN',
  4441. 'eSLDB',
  4442. 'iBKH',
  4443. 'miRTarBase', # Note: miRtarBase in the list
  4444. 'other_sources'
  4445. ]
  4446. # Create a mapp

all_figures-raidhani.ipynb at commit 6067db0, under MIT · at the source

Overview

Authors: Gaurav Ahuja1, Arushi Sharma2, Ankit Singh1, Shekhar Kedia3, Abhinav Sharma4, Vishakha Gautam5, Pranjal Sharma1, Aniket Khandelwal1, Sarayu Ramakrishna6, Ravi Muddashetty6, Kristine Freude7, Saveena Solanki5, Sonam Chauhan2, Suvendu Kumar5, Shiva Satija2, Subhadeep Duari5, Sakshi Arora5, Advik Gupta1, Raidhani Shome8, Debarka Sengupta8, Deepak nair6
  1. Indraprastha Institute of Information Technology-Delhi (IIIT-Delhi)
  2. Indraprastha Institute of Information Technology
  3. University of Cambridge
  4. Indraprastha Institute of Information Technology-Delhi, (IIIT-Delhi)
  5. Department of Computational Biology, Indraprastha Institute of Information Technology-Delhi (IIIT-Delhi), Okhla, Phase III, New Delhi, 110020, India
  6. Indian Institute of Science
  7. University of Copenhagen
  8. Indraprastha Institute of Information Technology Delhi
Dates: published online 13 March 2026
Type: Preprint
License: CC BY
Identifiers: DOI 10.21203/rs.3.rs-9060414/v1 · OpenAlex W7135197563
Open access: green, a free copy (OpenAlex)
Status: code verified
Categories: human (organism), Alzheimer's / dementia (population), cellular / molecular (subfield)
Methods: Statistics, Smoothing, state filtering, decompositions
Keywords: Knowledge Graphs, Evolution, Graph Modeling, Aging, Synapse, Alzheimer
Topic: Bioinformatics and Genomic Networks (Molecular Biology, Biochemistry, Genetics and Molecular Biology), according to OpenAlex
Citations: not cited yet (Europe PMC); 107 references in the paper

Abstract

Aging research has been advanced largely through the use of model organisms, where short lifespans and genetic tractability enable the systematic discovery of molecular pathways influencing longevity and age-related decline. However, knowledge about aging remains fragmented across species-specific repositories and domain-focused databases, limiting our ability to identify evolutionarily conserved mechanisms and translate findings to human biology. To address this gap, we developed EvoAge, a unified, multi-species knowledge graph that integrates aging-specific and general biomedical resources into a systems-level framework. EvoAge harmonizes 48 public datasets into a graph comprising 1.04 billion triples across six key species. A human-centric orthology framework reconciles more than 80,000 gene entries, expanding accessible organism-level aging knowledge by up to 1,700-fold compared with existing resources. To operationalize the graph for biological reasoning, we optimized knowledge graph embedding models and deployed a large language model (LLM)-assisted agentic interface that supports natural-language querying, link prediction, and hypothesis testing. In internal benchmarking using recent pre-print aging literature, EvoAge significantly outperformed state-of-the-art LLMs in distinguishing biologically plausible from implausible hypotheses. Importantly, EvoAge recommended a previously unrecognized Alzheimer’s disease (AD) mechanism involving nanoscale redistribution of BACE1 within synaptic compartments. We experimentally validated this EvoAge-supported prediction using patient-derived iPSCs carrying a familial PSEN1 mutation, demonstrating disease-associated remodeling of β-secretase, defined by altered localization, nanoscale clustering, and compartment-specific enrichment. We further confirmed the predicted evolutionary conservation of this BACE1–pathology relationship in additional AD systems, including transgenic mice and postmortem human brain tissue.

Reproduced under the paper's license (CC BY), from the paper cited above.

Repositories

Its files are read in the Code ↔ Paper reader above, with 21 matches between paragraphs and lines of code.

the-ahuja-lab/EvoAge

License: MIT
State: the link answers, verified on 30 September 2026
Evidence: files inventoried
Commit: 6067db0a3ac902952f48933db8e72dac4fd15901, 25 August 2026
Languages: Jupyter (221), Python (162), Shell (35), JavaScript (2), R (1)
Size: 2,385 files, 421 scripts
Software Heritage: not archived
Found in: “Data and Code Availability”
Holds: README, license file, environment (Backend/poetry.lock, Backend/pyproject.toml, Backend/requirements.txt, Frontend/requirements.txt, Frontend/evo-utils/pyproject.toml, Backend/dgl-ke/python/setup.cfg, Backend/dgl-ke/python/setup.py, pipeline/09_evoage_vs_other/evoage_vs_biochat_escargot/escargot/escargot/pyproject.toml), tests, continuous integration, documentation, 221 notebooks
Not found: CITATION.cff
Tools: pandas (255 files), NumPy (212 files), PyTorch (34 files), SciPy (7 files), Matplotlib (4 files), scikit-learn (2 files), CuPy (1 file), NetworkX (1 file), Pillow (1 file), Plotly (1 file), tidyverse (1 file)
Availability: 1 check, the latest on 30 September 2026: the link answers
  • 30 September 2026: the link answers
421 files

hub.docker.com/r/ahujalab

License: none: the authors keep all their rights
State: the link answers, verified on 30 September 2026
Evidence: the link answers
Software Heritage: not checked
Found in: “Data and Code Availability”
Not found: README, license file, CITATION.cff, environment file, tests, continuous integration, documentation
Availability: 1 check, the latest on 30 September 2026: the link answers (HTTP 200)
  • 30 September 2026: the link answers (HTTP 200)

The paper's code and data availability statement is in the Data section.

Tracing map

Proposed by the machine: these links were found in the paper and verified at the source, without human review. The map will receive a Zenodo DOI once one of the paper's authors has validated it with their ORCID.

What the map holds:

  • 2 repositories of the authors' code, each at its verified commit, with its license and how the link was found in the paper;
  • 419 scripts, each with its path and the digest of its content;
  • 21 matches between paragraphs of the paper and lines of the code (method lexical-v1);
  • neither the text of the paper nor the code itself.

Its JSON (tracing-map.json) is deposited on Zenodo with its DOI once the map is validated.

Data

Datasets cited

Data and Code Availability

The complete source code for the EvoAge platform is publicly available on GitHub (https://github.com/the-ahuja-lab/EvoAge). The EvoAge chatbot web server is publicly accessible at https://evoage.ahujalab.iiitd.edu.in/. The full dataset generated and used in this study, including the final knowledge graph, is archived on Zenodo at https://doi.org/10.5281/zenodo.17711173. For enhanced reproducibility, a pre-configured Docker container is also available on Docker Hub at https://hub.docker.com/r/ahujalab/evoage-project.

Reproduced under the paper's license (CC BY), from the paper cited above.

Versions

The history of this record: each version stored by the harvester or made by a correction of its authors or of the maintainers of its code, and what changed in its facts. The texts of the paper (its abstract, its availability statements) are not part of it; versions that changed only those are not listed.

Version 1, 30 September 2026: the first record

Recorded: type, journal, dates, 21 authors, 6 keywords, 84 references.

Cite

This paper

Ahuja, G., Sharma, A., Singh, A., Kedia, S., Sharma, A., Gautam, V., Sharma, P., Khandelwal, A., Ramakrishna, S., Muddashetty, R., Freude, K., Solanki, S., Chauhan, S., Kumar, S., Satija, S., Duari, S., Arora, S., Gupta, A., Shome, R., . . . nair, D. (2026). Cross-Species Aging Knowledge Integration into Agentic AI Platform Uncovers Conserved Mechanisms. Research Square (preprint). https://doi.org/10.21203/rs.3.rs-9060414/v1

BibTeX

@article{ahuja2026cross,
author = {Ahuja, Gaurav and Sharma, Arushi and Singh, Ankit and Kedia, Shekhar and Sharma, Abhinav and Gautam, Vishakha and Sharma, Pranjal and Khandelwal, Aniket and Ramakrishna, Sarayu and Muddashetty, Ravi and Freude, Kristine and Solanki, Saveena and Chauhan, Sonam and Kumar, Suvendu and Satija, Shiva and Duari, Subhadeep and Arora, Sakshi and Gupta, Advik and Shome, Raidhani and Sengupta, Debarka and nair, Deepak},
title = {{Cross-Species Aging Knowledge Integration into Agentic AI Platform Uncovers Conserved Mechanisms}},
journal = {Research Square (preprint)},
year = {2026},
month = mar,
publisher = {Research Square},
issn = {2693-5015},
doi = {10.21203/rs.3.rs-9060414/v1},
url = {https://doi.org/10.21203/rs.3.rs-9060414/v1}
}

RIS

TY - JOUR
AU - Ahuja, Gaurav
AU - Sharma, Arushi
AU - Singh, Ankit
AU - Kedia, Shekhar
AU - Sharma, Abhinav
AU - Gautam, Vishakha
AU - Sharma, Pranjal
AU - Khandelwal, Aniket
AU - Ramakrishna, Sarayu
AU - Muddashetty, Ravi
AU - Freude, Kristine
AU - Solanki, Saveena
AU - Chauhan, Sonam
AU - Kumar, Suvendu
AU - Satija, Shiva
AU - Duari, Subhadeep
AU - Arora, Sakshi
AU - Gupta, Advik
AU - Shome, Raidhani
AU - Sengupta, Debarka
AU - nair, Deepak
TI - Cross-Species Aging Knowledge Integration into Agentic AI Platform Uncovers Conserved Mechanisms
T2 - Research Square (preprint)
J2 - Res Sq
PY - 2026
DA - 2026/03/13
SN - 2693-5015
PB - Research Square
DO - 10.21203/rs.3.rs-9060414/v1
UR - https://doi.org/10.21203/rs.3.rs-9060414/v1
ER -

CSL-JSON

{
"id": "10.21203/rs.3.rs-9060414/v1",
"type": "article",
"title": "Cross-Species Aging Knowledge Integration into Agentic AI Platform Uncovers Conserved Mechanisms",
"container-title": "Research Square (preprint)",
"author": [
{
"family": "Ahuja",
"given": "Gaurav"
},
{
"family": "Sharma",
"given": "Arushi"
},
{
"family": "Singh",
"given": "Ankit"
},
{
"family": "Kedia",
"given": "Shekhar"
},
{
"family": "Sharma",
"given": "Abhinav"
},
{
"family": "Gautam",
"given": "Vishakha"
},
{
"family": "Sharma",
"given": "Pranjal"
},
{
"family": "Khandelwal",
"given": "Aniket"
},
{
"family": "Ramakrishna",
"given": "Sarayu"
},
{
"family": "Muddashetty",
"given": "Ravi"
},
{
"family": "Freude",
"given": "Kristine"
},
{
"family": "Solanki",
"given": "Saveena"
},
{
"family": "Chauhan",
"given": "Sonam"
},
{
"family": "Kumar",
"given": "Suvendu"
},
{
"family": "Satija",
"given": "Shiva"
},
{
"family": "Duari",
"given": "Subhadeep"
},
{
"family": "Arora",
"given": "Sakshi"
},
{
"family": "Gupta",
"given": "Advik"
},
{
"family": "Shome",
"given": "Raidhani"
},
{
"family": "Sengupta",
"given": "Debarka"
},
{
"family": "nair",
"given": "Deepak"
}
],
"container-title-short": "Res Sq",
"DOI": "10.21203/rs.3.rs-9060414/v1",
"ISSN": "2693-5015",
"publisher": "Research Square",
"URL": "https://doi.org/10.21203/rs.3.rs-9060414/v1",
"issued": {
"date-parts": [
[
2026,
3,
13
]
]
}
}

The tracing map gets a citation of its own once an author has validated it and it has a DOI.

Similar papers

The papers with a page that share the most with this one: the tools found in their code, their categories, datasets, cited references and authors, the rarest counting most.

[1] doi:10.1038/s42003-026-10957-8 [code]
Brain defence by the extracellular matrix protein Cochlin.
Journal: Communications biology
In common: CuPy, NetworkX, Plotly, 8 other tools, cellular / molecular
[2] doi:10.1016/j.isci.2026.116825 [code]
Social hierarchy shapes behavioral and transcriptional responses to chronic stress and ketamine in male mice.
Journal: iScience
In common: CuPy, NetworkX, Plotly, 7 other tools
[3] doi:10.1016/j.isci.2026.116055 [code]
Mapping the transcriptional diversity of calcium signaling in the mouse and human brain.
Journal: iScience
In common: CuPy, NetworkX, Plotly, 7 other tools
[4] doi:10.1162/imag.a.1276 [code]
High-resolution whole-brain magnetic resonance spectroscopic imaging in youth at risk for psychosis.
Journal: Imaging neuroscience (Cambridge, Mass.)
In common: CuPy, NetworkX, Plotly, 6 other tools
[5] doi:10.1186/s13293-026-00927-4 [code]
Gene regulatory network analysis identifies dysregulation of hypoxia pathways as contributing to glioblastoma treatment resistance in females.
Journal: Biology of sex differences
In common: CuPy, NetworkX, PyTorch, 5 other tools, cellular / molecular, 1 reference
[6] doi:10.1038/s41514-026-00391-9 [code]
Region-specific transcriptional signatures of brain aging in the absence of neuropathology at the single-cell level.
Journal: npj aging
In common: Plotly, PyTorch, tidyverse, 5 other tools, cellular / molecular, 2 references
[7] doi:10.1371/journal.pcbi.1014555 [code]
Body surface potential driven personalisation of electrophysiological digital twins in hypertrophic cardiomyopathy.
Journal: PLoS computational biology
In common: CuPy, NetworkX, Pillow, 6 other tools
[8] doi:10.64898/2026.03.30.714220 [code]
An integrated single cell and spatial omics atlas of human prenatal development
Journal: bioRxiv (preprint)
In common: CuPy, NetworkX, Pillow, 6 other tools
[9] doi:10.1038/s41593-026-02267-3 [code]
Spatial proteomic analysis in human Alzheimer's disease brains enables identification of microenvironment-dependent microglial cell states.
Journal: Nature neuroscience
In common: CuPy, NetworkX, Plotly, 5 other tools, Alzheimer's / dementia, cellular / molecular
[10] doi:10.1038/s41467-026-71525-6 [code]
Single-nucleus brain transcriptomics reveals microglia dysfunction in multiple system atrophy.
Journal: Nature communications
In common: tidyverse, pandas, Matplotlib, 1 other tool, cellular / molecular, 1 reference, author Kristine Freude

Contribute

The authors of this paper can claim it, correct its record and validate its tracing map, and the maintainers of its code (its owner, or a public member of its organization) correct what it says of their repository; anyone signed in can ask for its removal. Every request goes to OSCR's own machine, which answers it; your account page follows them.

Sign in with ORCID to claim this paper as one of its authors, correct its record or validate its tracing map: when the paper's metadata lists your ORCID iD, you are recognized at once. Maintainers of its code: sign in with GitHub, then claim the repository on your account page.

Request its removal

To ask OSCR to remove this record, the copies of its authors' scripts or its tracing map, use the removal request page: signed in, you say who you are, what to remove and why, then review and confirm the request. Published rules decide every request (how).

Discussion, reproductions, activity

Discussion: questions and error reports about this paper and its code, from signed-in readers and its authors. It opens with sign-in.

Reproductions: reports from readers who ran the authors' code: what they reproduced, with which environment, commit and data. It opens with sign-in.

Activity: what happens around this paper: new versions of its record, its map's validation, discussions and reproductions. It opens with sign-in.