from matplotlib.lines import Line2D
try:
from figure_style import apply_figure_style, panel_letter, META_GREY # if skill kernel present
apply_figure_style()
except Exception:
META_GREY = "#8a8a8a"
def panel_letter(ax, s): ax.text(-0.08, 1.04, s, transform=ax.transAxes, fontweight="bold", fontsize=11)
heat = fst_df.set_index(fst_df["gene"] + " " + fst_df["rsid"])[[f"AF_{sp}" for sp in SUPER]].astype(float)
heat.columns = SUPER
row_order = ["MFSD12 rs2240751","MFSD12 rs10424065","BNC2 rs16935073","BNC2 rs2153271",
"SPIRE2 rs12598316","SPIRE2 rs34357723","TSPAN10 rs6420484"]
heat = heat.loc[row_order]
ital = lambda r: (r.replace("MFSD12","$\\it{MFSD12}$").replace("BNC2","$\\it{BNC2}$")
.replace("SPIRE2","$\\it{SPIRE2}$").replace("TSPAN10","$\\it{TSPAN10}$"))
ylabs = [ital(r) for r in heat.index]
gene_col = {"MFSD12":"#c0392b","SPIRE2":"#2980b9","BNC2":"#27ae60","TSPAN10":"#8e44ad"}
med, p95 = np.median(baseline_arr), np.percentile(baseline_arr,95)
fig = plt.figure(figsize=(11.5,4.6))
gs = fig.add_gridspec(1,2, width_ratios=[1.5,1.15], wspace=0.55)
axA, axB = fig.add_subplot(gs[0]), fig.add_subplot(gs[1])
im = axA.imshow(heat.values, cmap="viridis", aspect="auto", vmin=0, vmax=heat.values.max())
axA.set_xticks(range(5)); axA.set_xticklabels(SUPER)
axA.set_yticks(range(len(heat))); axA.set_yticklabels(ylabs)
for i in range(len(heat)):
for j in range(5):
v = heat.values[i,j]
axA.text(j,i,f"{v:.2f}",ha="center",va="center",
color="white" if v<0.55*heat.values.max() else "black",fontsize=6)
axA.set_title("Associated-allele frequency across 1000G superpopulations",fontsize=8)
cb = fig.colorbar(im, ax=axA, fraction=0.046, pad=0.03); cb.set_label("assoc. allele freq",fontsize=6)
for i0,i1 in [(0,1),(2,3),(4,5)]:
axA.plot([-0.72,-0.72],[i0-0.4,i1+0.4],color="#c0392b",lw=1.4,clip_on=False)
axA.annotate("mirror\npairs", xy=(-0.72,0.5), xytext=(-2.0,0.5), rotation=90,
va="center", ha="center", fontsize=6, color="#c0392b", annotation_clip=False)
d = fst_df.sort_values("fst_hudson_5superpop").reset_index(drop=True)
axB.axvspan(baseline_arr.min(), p95, color=META_GREY, alpha=0.12, zorder=0)
axB.axvline(med, color=META_GREY, lw=1, zorder=1)
axB.axvline(p95, color=META_GREY, lw=0.8, ls="--", zorder=1)
for y,(_,rr) in enumerate(d.iterrows()):
axB.plot([0,rr["fst_hudson_5superpop"]],[y,y],color=gene_col[rr["gene"]],lw=1,alpha=0.6,zorder=2)
axB.scatter(rr["fst_hudson_5superpop"],y,s=46,color=gene_col[rr["gene"]],edgecolor="black",lw=0.5,zorder=3)
axB.text(rr["fst_hudson_5superpop"]+0.007,y,f"p{rr['baseline_percentile']:.0f}",va="center",fontsize=5.5,color="#333")
axB.set_yticks(range(len(d)))
axB.set_yticklabels([f"$\\it{{{r.gene}}}$ {r.rsid}" for r in d.itertuples()], fontsize=6)
axB.set_xlabel("Hudson $F_{ST}$ (5 superpops)")
axB.set_title("Convergent variants vs genome-wide background", fontsize=8, pad=14)
axB.text(med, len(d)-0.5, "baseline median", fontsize=5, color="#555", ha="center")
axB.text(p95+0.004, len(d)-0.5, "95th pct", fontsize=5, color="#555", ha="left")
axB.set_xlim(0, d["fst_hudson_5superpop"].max()*1.2); axB.set_ylim(-0.7, len(d)-0.3)
axB.legend(handles=[Line2D([0],[0],marker="o",color="w",markerfacecolor=c,markeredgecolor="k",
markersize=6,label=g) for g,c in gene_col.items()],
fontsize=5.5, frameon=False, loc="lower right", title="gene", title_fontsize=6)
panel_letter(axA,"a"); panel_letter(axB,"b")
os.makedirs(FIGDIR, exist_ok=True)
fig.savefig(os.path.join(FIGDIR,"nb11_cross_ancestry.png"), dpi=300, bbox_inches="tight")
plt.show()