Back to home page

EIC code displayed by LXR

 
 

    


File indexing completed on 2026-07-26 08:22:17

0001 import pandas
0002 import matplotlib.pyplot
0003 import pathlib
0004 import numpy
0005 import scipy
0006 
0007 # A list of plots, described as tuples of (category, name, type), where
0008 # category is any of "seeding", "fitting", and "finding"; name is a tree name
0009 # in the output ROOT file, and type is any of "eff", "prod", or "hist".
0010 PLOT_NAMES = [
0011     ("seeding", "seeding_trackeff_vs_pT", "eff"),
0012     ("seeding", "seeding_trackeff_vs_eta", "eff"),
0013     ("seeding", "seeding_trackeff_vs_phi", "eff"),
0014     ("finding", "finding_trackeff_vs_pT", "eff"),
0015     ("finding", "finding_trackeff_vs_eta", "eff"),
0016     ("finding", "finding_trackeff_vs_phi", "eff"),
0017     ("finding", "finding_nDuplicated_vs_eta", "prof"),
0018     ("finding", "finding_nFakeTracks_vs_eta", "prof"),
0019     ("finding", "ndf", "hist"),
0020     ("finding", "pval", "hist"),
0021     ("finding", "purity", "hist"),
0022     ("finding", "completeness", "hist"),
0023     ("fitting", "res_d0", "hist"),
0024     ("fitting", "res_z0", "hist"),
0025     ("fitting", "res_phi", "hist"),
0026     ("fitting", "res_qop", "hist"),
0027     ("fitting", "res_qopT", "hist"),
0028     ("fitting", "res_qopz", "hist"),
0029     ("fitting", "res_theta", "hist"),
0030     ("fitting", "pull_d0", "hist"),
0031     ("fitting", "pull_z0", "hist"),
0032     ("fitting", "pull_phi", "hist"),
0033     ("fitting", "pull_qop", "hist"),
0034     ("fitting", "pull_theta", "hist"),
0035     ("fitting", "ndf", "hist"),
0036     ("fitting", "pval", "hist"),
0037 ]
0038 
0039 # Ratio plots encoded as tuples of two plots, where each individual plot is of
0040 # the form (category, name) as described above.
0041 RATIO_PLOTS = [
0042     (
0043         ("seeding", "seeding_trackeff_vs_pT"),
0044         ("finding", "finding_trackeff_vs_pT"),
0045     ),
0046     (
0047         ("seeding", "seeding_trackeff_vs_eta"),
0048         ("finding", "finding_trackeff_vs_eta"),
0049     ),
0050     (
0051         ("seeding", "seeding_trackeff_vs_phi"),
0052         ("finding", "finding_trackeff_vs_phi"),
0053     ),
0054 ]
0055 
0056 
0057 def make_plots(plot_candidate_names, output_dir, fig_kwargs=None, file_base=None):
0058     color_spacing = max(len(plot_candidate_names.keys()), 2)
0059 
0060     files_to_copy = []
0061 
0062     if fig_kwargs is None:
0063         fig_kwargs = {}
0064 
0065     if file_base is None:
0066         file_base = ""
0067     else:
0068         file_base = file_base + "_"
0069 
0070     for data_cat, plot, plot_type in PLOT_NAMES:
0071         dfs = {
0072             k: pandas.read_csv(d / data_cat / (plot + ".csv"))
0073             for (k, (d, _)) in plot_candidate_names.items()
0074         }
0075 
0076         fig = matplotlib.pyplot.figure(**fig_kwargs)
0077         ax = fig.subplots()
0078 
0079         for k, v in dfs.items():
0080             v["center"] = (v["bin_left"] + v["bin_right"]) / 2
0081             v["width"] = v["center"] - v["bin_left"]
0082             if "trackeff" in plot:
0083                 dfs[k] = v[v["ntotal"] > 0]
0084 
0085         if plot_type == "eff":
0086             ax.set_ylabel("Efficiency")
0087         elif plot == "finding_nDuplicated_vs_eta":
0088             ax.set_ylabel("Duplicate rate")
0089         elif plot == "finding_nFakeTracks_vs_eta":
0090             ax.set_ylabel("Fake rate")
0091         else:
0092             ax.set_ylabel("Normalized entries")
0093 
0094         plot_normal = False
0095 
0096         if "vs_eta" in plot:
0097             ax.set_xlabel("$\\eta$")
0098         elif "vs_phi" in plot:
0099             ax.set_xlabel("$\\phi$")
0100         elif "vs_pT" in plot:
0101             ax.set_xlabel("$p_T$ (GeV)")
0102         elif "pval" in plot:
0103             ax.set_xlabel("$p$")
0104         elif "ndf" in plot:
0105             ax.set_xlabel("NDF")
0106         elif "completeness" in plot:
0107             ax.set_xlabel("Completeness")
0108         elif "purity" in plot:
0109             ax.set_xlabel("Purity")
0110         elif "res_d0" == plot:
0111             ax.set_xlabel("Residual $d_0$")
0112         elif "res_z0" == plot:
0113             ax.set_xlabel("Residual $z_0$")
0114         elif "res_phi" == plot:
0115             ax.set_xlabel("Residual $\\phi$")
0116         elif "res_qop" == plot:
0117             ax.set_xlabel("Residual $q/p$")
0118         elif "res_qopT" == plot:
0119             ax.set_xlabel("Residual $q/p_T$")
0120         elif "res_theta" == plot:
0121             ax.set_xlabel("Residual $\\theta$")
0122         elif "res_qopz" == plot:
0123             ax.set_xlabel("Residual $q/p_z$")
0124         elif "pull_d0" == plot:
0125             plot_normal = True
0126             ax.set_xlabel("Pull $d_0$")
0127         elif "pull_z0" == plot:
0128             plot_normal = True
0129             ax.set_xlabel("Pull $z_0$")
0130         elif "pull_phi" == plot:
0131             plot_normal = True
0132             ax.set_xlabel("Pull $\\phi$")
0133         elif "pull_qop" == plot:
0134             plot_normal = True
0135             ax.set_xlabel("Pull $q/p$")
0136         elif "pull_theta" == plot:
0137             plot_normal = True
0138             ax.set_xlabel("Pull $\\theta$")
0139 
0140         if data_cat == "seeding":
0141             base_color = 0 * color_spacing
0142         elif data_cat == "finding":
0143             base_color = 1 * color_spacing
0144         else:
0145             base_color = 2 * color_spacing
0146 
0147         kwargs = {}
0148         plot_kwargs = {k: {} for k in dfs.keys()}
0149         scale_factors = {k: 1 for k in dfs.keys()}
0150 
0151         if plot_type == "hist" or plot_type == "prof":
0152             kwargs["drawstyle"] = "steps-mid"
0153         else:
0154             kwargs["fmt"] = "."
0155             for k in dfs.keys():
0156                 plot_kwargs[k]["xerr"] = dfs[k]["width"]
0157 
0158         if plot_type == "eff":
0159             xkey = "efficiency"
0160         elif plot_type == "prof":
0161             xkey = "value"
0162         elif plot_type == "hist":
0163             xkey = "ntotal"
0164             for k in dfs.keys():
0165                 scale_factors[k] = dfs[k][xkey].sum()
0166 
0167         labels = {k: plot_candidate_names[k][1] for k in dfs.keys()}
0168 
0169         if "pull_" in plot or "res_" in plot:
0170             for k in dfs.keys():
0171                 mean = numpy.average(dfs[k]["center"], weights=dfs[k][xkey])
0172                 std = numpy.sqrt(
0173                     numpy.average((dfs[k]["center"] - mean) ** 2, weights=dfs[k][xkey])
0174                 )
0175                 labels[k] = labels[k] + "; $\\mu = %.3f$, $\\sigma = %.3f$" % (
0176                     mean,
0177                     std,
0178                 )
0179 
0180         if plot_normal:
0181             x = numpy.linspace(-5, 5, 200)
0182             ax.plot(x, scipy.stats.norm.pdf(x, 0, 1), label="Ideal", color="black")
0183 
0184         if plot_type == "hist":
0185             for k in dfs.keys():
0186                 scale_factors[k] *= (dfs[k]["bin_right"] - dfs[k]["bin_left"])[0]
0187 
0188         for i, k in enumerate(dfs.keys()):
0189             ax.errorbar(
0190                 dfs[k]["center"],
0191                 dfs[k][xkey] / scale_factors[k],
0192                 yerr=(
0193                     dfs[k]["err_low"] / scale_factors[k],
0194                     dfs[k]["err_high"] / scale_factors[k],
0195                 ),
0196                 capsize=3,
0197                 label=labels[k],
0198                 color="C%d" % (base_color + i),
0199                 **kwargs,
0200                 **plot_kwargs[k],
0201             )
0202 
0203         first_key = list(dfs.keys())[0]
0204         ax.set_xlim(
0205             xmin=dfs[first_key]["bin_left"].min(),
0206             xmax=dfs[first_key]["bin_right"].max(),
0207         )
0208 
0209         ax.legend()
0210         fig.tight_layout()
0211         n_data_cat = data_cat
0212         fig.savefig(
0213             pathlib.Path(output_dir) / ("%s%s_%s.png" % (file_base, n_data_cat, plot))
0214         )
0215 
0216         matplotlib.pyplot.close()
0217 
0218         files_to_copy.append(
0219             "%s/%s%s_%s.png" % (output_dir, file_base, n_data_cat, plot)
0220         )
0221 
0222     for (from_data_cat, from_plot), (to_data_cat, to_plot) in RATIO_PLOTS:
0223         dfs = {
0224             k: (
0225                 pandas.read_csv(d / from_data_cat / (from_plot + ".csv")),
0226                 pandas.read_csv(d / to_data_cat / (to_plot + ".csv")),
0227             )
0228             for (k, (d, _)) in plot_candidate_names.items()
0229         }
0230 
0231         fig = matplotlib.pyplot.figure(**fig_kwargs)
0232         ax = fig.subplots()
0233 
0234         masks = {}
0235         ratios = {}
0236         yerr_lows = {}
0237         yerr_highs = {}
0238 
0239         for k in dfs.keys():
0240             dfs[k][0]["center"] = (dfs[k][0]["bin_left"] + dfs[k][0]["bin_right"]) / 2
0241             dfs[k][0]["width"] = dfs[k][0]["center"] - dfs[k][0]["bin_left"]
0242 
0243             masks[k] = dfs[k][0]["ntotal"] > 0
0244             ratios[k] = (
0245                 dfs[k][1][masks[k]]["efficiency"] / dfs[k][0][masks[k]]["efficiency"]
0246             )
0247 
0248             yerr_lows[k] = ratios[k] * numpy.sqrt(
0249                 (dfs[k][0][masks[k]]["err_low"] / dfs[k][0][masks[k]]["efficiency"])
0250                 ** 2
0251                 + (dfs[k][1][masks[k]]["err_low"] / dfs[k][1][masks[k]]["efficiency"])
0252                 ** 2
0253             )
0254             yerr_highs[k] = ratios[k] * numpy.sqrt(
0255                 (dfs[k][0][masks[k]]["err_high"] / dfs[k][0][masks[k]]["efficiency"])
0256                 ** 2
0257                 + (dfs[k][1][masks[k]]["err_high"] / dfs[k][1][masks[k]]["efficiency"])
0258                 ** 2
0259             )
0260 
0261         ax.set_ylabel("Relative efficiency")
0262 
0263         if "vs_eta" in from_plot:
0264             ax.set_xlabel("$\\eta$")
0265         elif "vs_phi" in from_plot:
0266             ax.set_xlabel("$\\phi$")
0267         elif "vs_pT" in from_plot:
0268             ax.set_xlabel("$p_T$ (GeV)")
0269 
0270         if from_data_cat == "seeding":
0271             base_color = 3 * color_spacing
0272         elif from_data_cat == "finding":
0273             base_color = 4 * color_spacing
0274         else:
0275             base_color = 5 * color_spacing
0276 
0277         for i, k in enumerate(dfs.keys()):
0278             ax.errorbar(
0279                 dfs[k][0][masks[k]]["center"],
0280                 ratios[k],
0281                 yerr=(yerr_lows[k], yerr_highs[k]),
0282                 capsize=3,
0283                 label=plot_candidate_names[k][1],
0284                 color="C%d" % (base_color + i),
0285                 fmt=".",
0286                 xerr=dfs[k][0][masks[k]]["width"],
0287             )
0288 
0289         first_key = list(dfs.keys())[0]
0290         ax.set_xlim(
0291             xmin=dfs[first_key][0][masks[first_key]]["bin_left"].min(),
0292             xmax=dfs[first_key][0][masks[first_key]]["bin_right"].max(),
0293         )
0294 
0295         ax.legend()
0296         fig.tight_layout()
0297 
0298         vs_name = "%svs%s" % (from_data_cat, to_data_cat)
0299 
0300         fig.savefig(
0301             pathlib.Path(output_dir) / ("%s%s_%s.png" % (file_base, vs_name, from_plot))
0302         )
0303         matplotlib.pyplot.close()
0304 
0305         files_to_copy.append(
0306             "%s/%s%s_%s.png" % (output_dir, file_base, n_data_cat, plot)
0307         )
0308 
0309     return files_to_copy