from matplotlib.backends.backend_pdf import PdfPages
atlas_index=[]
atlas_page_counts={}
for group in ("sst","wt"):
    rows=R[R.group.eq(group)&R.primary].sort_values("unit_id")
    levels=sorted(DATA[group]["trials"].level.unique())
    with PdfPages(OUT/(group+"-all-unit-rasters.pdf")) as pdf:
        for start in range(0,len(rows),8):
            page=rows.iloc[start:start+8]
            fig,axes=plt.subplots(8,4,figsize=(15,14))
            fig.subplots_adjust(left=.055,right=.985,bottom=.045,top=.94,wspace=.24,hspace=.62)
            for i,row in enumerate(page.itertuples()):
                levelrows=R[R.group.eq(group)&R.unit_id.eq(row.unit_id)].set_index("level")
                for j in range(4):
                    ax=axes[i,j];lev=row.level if j==0 else levels[j-1]
                    raster(ax,group,row.unit_id,lev,wide=(j==0))
                    if j==0:
                        ax.set_title(f"{row.unit_id} | {row.region}",loc="left",fontsize=8)
                        ax.set_ylabel("Trial")
                    else:
                        r=levelrows.loc[lev]
                        med=f"{r.first_spike_median_ms:.2f}" if np.isfinite(r.first_spike_median_ms) else "n.d."
                        status=("response" if bool(r.responder) else "not classified") if group=="sst" else "descriptive"
                        ax.set_title(f"Level {lev:g}; median {med} ms; {status}",loc="left",fontsize=7)
                        if np.isfinite(r.onset_ms):ax.axvline(r.onset_ms,color="black",ls="--",lw=.7)
                    if i==len(page)-1:ax.set_xlabel("Time from onset (ms)")
                atlas_index.append({"group":group,"unit_id":row.unit_id,"page":start//8+1,"row":i+1,"primary_onset_ms":row.onset_ms})
            for i in range(len(page),8):
                for ax in axes[i]:ax.set_visible(False)
            fig.suptitle(f"{group.upper()} trial atlas | page {start//8+1} | left: primary level, wide view; right: all levels, detail\nBlue shading: 0–10 ms light; gray: excluded onset/offset intervals; dashed line: resolved onset; median: conditional first spike in [2,8) ms",fontsize=8,y=.979)
            fig.canvas.draw()
            pdf.savefig(fig,dpi=100)
            if start==0:fig.savefig(OUT/(group+"-atlas-first-page.png"),dpi=100)
            plt.close(fig)
    atlas_page_counts[group]=(len(rows)+7)//8
pd.DataFrame(atlas_index).to_csv(OUT/"raster-atlas-index.csv",index=False)
print("Atlas pages",atlas_page_counts,"unit rows",len(atlas_index))
