from __future__ import annotations
import functools as ft
from typing import Literal
import itertools as it
from .style import StyleSet
from analyzer.utils.structure_tools import (
ItemWithMeta,
commonDict,
dictToDot,
dotFormat,
)
from .processors import BasePostprocessor
from .plots.plots_1d import plotOne, plotRatio, plotRatioOfRatios, plotModel
from .plots.plots_2d import plot2D
from attrs import define, field
[docs]
ResultSet = list[list[ItemWithMeta]]
@define
[docs]
class Histogram1D(BasePostprocessor):
[docs]
style_set: str | StyleSet = field(factory=StyleSet)
[docs]
scale: Literal["log", "linear"] = "linear"
[docs]
normalize: bool = False
[docs]
show_stacked_unc: bool = True
[docs]
def getRunFuncs(self, group, prefix=None):
if isinstance(group, dict):
unstacked = group["unstacked"]
stacked = group["stacked"]
else:
unstacked = group
stacked = None
common_meta = commonDict(it.chain((stacked or []), (unstacked or [])))
output_path = dotFormat(
self.output_name, **dict(dictToDot(common_meta)), prefix=prefix
)
pc = self.plot_configuration.makeFormatted(common_meta)
yield ft.partial(
plotOne,
unstacked,
stacked,
common_meta,
output_path,
scale=self.scale,
style_set=self.style_set,
normalize=self.normalize,
plot_configuration=pc,
show_stacked_unc=self.show_stacked_unc,
)
@define
[docs]
class RatioPlot(BasePostprocessor):
[docs]
scale: Literal["log", "linear"] = "linear"
[docs]
normalize: bool = False
[docs]
ratio_ylim: tuple[float, float] = (0, 2)
[docs]
ratio_hlines: list[float] = field(factory=lambda: [1.0])
[docs]
ratio_height: float = 0.5
[docs]
ratio_type: Literal["poisson", "poisson-ratio", "efficiency", "significance"] = (
"poisson"
)
[docs]
def getRunFuncs(self, group, prefix=None):
numerator = group["numerator"]
denominator = group["denominator"]
common_meta = commonDict(it.chain(numerator, denominator))
output_path = dotFormat(
self.output_name, prefix=prefix, **dict(dictToDot(common_meta))
)
pc = self.plot_configuration.makeFormatted(common_meta)
yield ft.partial(
plotRatio,
denominator,
numerator,
output_path,
self.style_set,
normalize=self.normalize,
ratio_ylim=self.ratio_ylim,
ratio_type=self.ratio_type,
scale=self.scale,
ratio_hlines=self.ratio_hlines,
ratio_height=self.ratio_height,
no_stack=self.no_stack,
plot_configuration=pc,
)
@define
[docs]
class RatioOfRatiosPlot(BasePostprocessor):
[docs]
r1_label: str = "{numerator.dataset_title}/{denominator.dataset_title}"
[docs]
r2_label: str = "{numerator.dataset_title}/{denominator.dataset_title}"
[docs]
double_ratio_label: str = "Double Ratio"
[docs]
scale: Literal["log", "linear"] = "linear"
[docs]
normalize: bool = False
[docs]
ratio_ylim: tuple[float, float] = (0, 2)
[docs]
ratio_hlines: list[float] = field(factory=lambda: [1.0])
[docs]
ratio_height: float = 0.5
[docs]
ratio_type: Literal["poisson", "poisson-ratio", "efficiency", "significance"] = (
"poisson"
)
[docs]
def getRunFuncs(self, group, prefix=None):
num_group = group["numerator"]
den_group = group["denominator"]
# Ensure we have exactly one item per component
if len(num_group["numerator"]) != 1:
raise ValueError(
f"Expected exactly 1 item for num_group['numerator'], got {len(num_group['numerator'])}"
)
if len(num_group["denominator"]) != 1:
raise ValueError(
f"Expected exactly 1 item for num_group['denominator'], got {len(num_group['denominator'])}"
)
if len(den_group["numerator"]) != 1:
raise ValueError(
f"Expected exactly 1 item for den_group['numerator'], got {len(den_group['numerator'])}"
)
if len(den_group["denominator"]) != 1:
raise ValueError(
f"Expected exactly 1 item for den_group['denominator'], got {len(den_group['denominator'])}"
)
num_numerator = num_group["numerator"][0]
num_denominator = num_group["denominator"][0]
den_numerator = den_group["numerator"][0]
den_denominator = den_group["denominator"][0]
common_meta = commonDict(
[num_numerator, num_denominator, den_numerator, den_denominator]
)
r1_meta = commonDict([num_numerator, num_denominator])
r2_meta = commonDict([den_numerator, den_denominator])
r1_label = dotFormat(self.r1_label, **dict(dictToDot(r1_meta)))
r2_label = dotFormat(self.r2_label, **dict(dictToDot(r2_meta)))
output_path = dotFormat(
self.output_name, prefix=prefix, **dict(dictToDot(common_meta))
)
pc = self.plot_configuration.makeFormatted(common_meta)
yield ft.partial(
plotRatioOfRatios,
num_numerator,
num_denominator,
den_numerator,
den_denominator,
output_path,
self.style_set,
r1_label=r1_label,
r2_label=r2_label,
double_ratio_label=self.double_ratio_label,
normalize=self.normalize,
ratio_ylim=self.ratio_ylim,
ratio_type=self.ratio_type,
scale=self.scale,
ratio_hlines=self.ratio_hlines,
ratio_height=self.ratio_height,
plot_configuration=pc,
)
@define
[docs]
class Histogram2D(BasePostprocessor):
[docs]
scale: Literal["log", "linear"] = "linear"
[docs]
normalize: bool = False
[docs]
def getRunFuncs(self, group, prefix=None):
if len(group) != 1:
raise RuntimeError()
hist = group[0]
common_meta = commonDict(group)
output_path = dotFormat(
self.output_name, prefix=prefix, **dict(dictToDot(common_meta))
)
self.plot_configuration.makeFormatted(common_meta)
yield ft.partial(
plot2D,
hist,
common_meta,
output_path,
style_set=self.style_set,
normalize=self.normalize,
plot_configuration=self.plot_configuration,
color_scale=self.scale,
)
@define
[docs]
class ModelPlot(BasePostprocessor):
[docs]
scale: Literal["log", "linear"] = "linear"
[docs]
normalize: bool = False
[docs]
ratio_ylim: tuple[float, float] = (0, 2)
[docs]
ratio_hlines: list[float] = field(factory=lambda: [1.0])
[docs]
ratio_height: float = 0.5
[docs]
ratio_type: Literal["poisson", "poisson-ratio", "efficiency", "significance"] = (
"poisson"
)
[docs]
def getRunFuncs(self, group, prefix=None):
data = group.get("data", [])
backgrounds = group.get("background", [])
signals = group.get("signal", []) or group.get("signals", [])
items = list(it.chain(data, backgrounds, signals))
if not items:
return
if len(data) != 1:
raise RuntimeError(f"Expected 1 data histogram, got {len(data)}")
data = data[0]
if not backgrounds:
raise RuntimeError("Expected at least 1 background histogram")
if len(signals) != 1:
raise RuntimeError(f"Expected 1 signal histogram, got {len(signals)}")
signal = signals[0]
common_meta = commonDict(items)
output_path = dotFormat(
self.output_name, prefix=prefix, **dict(dictToDot(common_meta))
)
pc = self.plot_configuration.makeFormatted(common_meta)
yield ft.partial(
plotModel,
data,
backgrounds,
signal,
output_path,
self.style_set,
normalize=self.normalize,
ratio_ylim=self.ratio_ylim,
ratio_type=self.ratio_type,
scale=self.scale,
ratio_hlines=self.ratio_hlines,
ratio_height=self.ratio_height,
plot_configuration=pc,
)