from attrs import define
from analyzer.core.columns import Column, TrackedColumns, EventBackend
import numpy as np
import awkward as ak
from analyzer.core.results import Histogram, UnscaledHistogram
import functools as ft
import operator as op
from analyzer.core.analysis_modules import (
AnalyzerModule,
MetadataExpr,
ModuleAddition,
PureResultModule,
)
from .axis import Axis, RegularAxis
import hist.dask as dah
import hist
import logging
[docs]
logger = logging.getLogger("analyzer.modules")
@define
[docs]
class HistogramBuilder(PureResultModule):
[docs]
storage: str = "weight"
[docs]
mask_col: Column | None = None
[docs]
def __attrs_post_init__(self):
if len(self.axes) == 0:
raise ValueError(f"HistogramBuilder '{self.product_name}' must have at least one axis defined.")
if len(self.axes) != len(self.columns):
raise ValueError(f"HistogramBuilder '{self.product_name}' has {len(self.axes)} axes but {len(self.columns)} columns.")
@staticmethod
@staticmethod
[docs]
def maybeFlatten(data):
if data.ndim == 2:
return ak.flatten(data)
else:
return data
@staticmethod
[docs]
def fillHistogram(
histogram,
cat_values,
fill_data,
weight=None,
variation="central",
mask=None,
):
all_values = (
[variation]
+ cat_values
+ [HistogramBuilder.maybeFlatten(x) for x in fill_data]
)
if weight is not None:
histogram.fill(*all_values, weight=weight)
else:
histogram.fill(*all_values)
return histogram
@staticmethod
[docs]
def create(backend, categories, axes, storage):
variations_axis = hist.axis.StrCategory([], name="variation", growth=True)
all_axes = (
[variations_axis]
+ [x.axis.toHist() for x in categories]
+ [x.toHist() for x in axes]
)
if backend == EventBackend.coffea_dask:
histogram = dah.Hist(*all_axes, storage=storage)
else:
histogram = hist.Hist(*all_axes, storage=storage)
return histogram
[docs]
def run(self, column_sets, params):
if isinstance(column_sets, TrackedColumns):
column_sets = [["central", column_sets]]
backend = column_sets[0][1].backend
pipeline_data = column_sets[0][1].pipeline_data
categories = pipeline_data.get("categories", {})
histogram = HistogramBuilder.create(
backend, categories, self.axes, self.storage
)
logger.debug(
f"Creating histogram {self.product_name} with the following variations:\n{[x[0] for x in column_sets]}"
)
to_run = [(x, y) for x, y in column_sets if x != "UNSCALED"]
for name, columns in to_run:
mask = None
if self.mask_col is not None:
mask = columns[self.mask_col]
data_to_fill = [columns[x] for x in self.columns]
if self.mask_col is not None:
data_to_fill = [col[mask] for col in data_to_fill]
representative = data_to_fill[0]
if "Weights" in columns.fields:
weights = columns["Weights"]
wf = weights.fields
if not wf:
total_weight = None
else:
wf = iter(wf)
total_weight = weights[next(wf)]
for w in wf:
total_weight = total_weight * weights[w]
total_weight = HistogramBuilder.transformToFill(
representative, total_weight, mask
)
else:
total_weight = None
cat_to_fill = [
HistogramBuilder.transformToFill(
representative, columns[x.column], mask
)
for x in categories
]
HistogramBuilder.fillHistogram(
histogram,
cat_to_fill,
data_to_fill,
weight=total_weight,
variation=name,
)
ret = [Histogram(name=self.product_name, histogram=histogram, axes=self.axes)]
if to_run := next((x for x in column_sets if x[0] == "UNSCALED"), None):
_, columns = to_run
mask = None
unscaled_histogram = HistogramBuilder.create(
backend, categories, self.axes, "double"
)
if self.mask_col is not None:
mask = columns[self.mask_col]
data_to_fill = [columns[x] for x in self.columns]
if self.mask_col is not None:
data_to_fill = [col[mask] for col in data_to_fill]
representative = data_to_fill[0]
cat_to_fill = [
HistogramBuilder.transformToFill(
representative, columns[x.column], mask
)
for x in categories
]
HistogramBuilder.fillHistogram(
unscaled_histogram,
cat_to_fill,
data_to_fill,
weight=None,
)
ret.append(
UnscaledHistogram(
name=self.product_name + "_unscaled",
histogram=unscaled_histogram,
axes=self.axes,
)
)
return ret
[docs]
def outputs(self, metadata):
return []
[docs]
def makeHistogram(
product_name: str,
columns,
axes: Axis | list[Axis],
data,
description=None,
mask=None,
):
"""
Create a histogram from column data and register it in the pipeline.
This helper function wraps input data into temporary columns, associates
them with axes definitions, and builds a histogram that can be added to
the module outputs.
Parameters
----------
product_name : str
Name of the histogram/product to create.
columns : list[Column]
Collection of event data columns where temporary columns will be stored.
axes : Axis or list of Axis
Axis (or axes) definition(s) for the histogram. Must match the dimensionality of `data`.
data : array-like or list of array-like
Data array(s) to histogram. Can be a single array or a list/tuple for multiple dimensions.
description : str, optional
Optional description for the histogram.
mask : array-like, optional
Boolean array indicating which entries should be included in the histogram.
Returns
-------
ModuleAddition
Object encapsulating the histogram builder, ready to be added to
an analyzer module's outputs.
Notes
-----
- Temporary columns are created for each data array to integrate with the pipeline.
- If `mask` is provided, it is stored in a separate column and used by the histogram builder.
"""
if not isinstance(data, (list, tuple)):
data = [data]
axes = [axes]
names = []
for i, d in enumerate(data):
name = Column(f"INTERNAL_USE.auto-col-{product_name}-{i}")
names.append(name)
columns[name] = d
if mask is not None:
mask_col_name = Column(f"INTERNAL_USE.mask-{product_name}")
columns[mask_col_name] = mask
else:
mask_col_name = None
b = HistogramBuilder(product_name, names, axes, mask_col=mask_col_name)
return ModuleAddition(b)
@define
[docs]
class SimpleHistogram(AnalyzerModule):
[docs]
axes: list[RegularAxis]
[docs]
replace_none: float | None = None
[docs]
mask_cols: list[Column] | None = None
[docs]
def __attrs_post_init__(self):
if len(self.axes) == 0:
raise ValueError(f"SimpleHistogram '{self.hist_name}' must have at least one axis defined.")
if len(self.axes) != len(self.input_cols):
raise ValueError(f"SimpleHistogram '{self.hist_name}' has {len(self.axes)} axes but {len(self.input_cols)} input columns.")
[docs]
def lint(self):
from analyzer.core.linting import LintLevel, LintMessage
total_bins = 1
for axis in self.axes:
total_bins *= axis.toHist().size
if total_bins > 2000:
return [
LintMessage(
level=LintLevel.WARNING,
category="HistogramDefinition",
message=f"Histogram '{self.hist_name}' has a total of {total_bins} bins across all axes. This may cause memory issues.",
module_name=self.name(),
)
]
return []
[docs]
def outputs(self, metadata):
return []
[docs]
def run(self, columns, params):
data = [columns[x] for x in self.input_cols]
if self.replace_none is not None:
data = [ak.fill_none(x, self.replace_none) for x in data]
if self.mask_cols is not None:
mask = ft.reduce(op.and_, [columns[x] for x in self.mask_cols])
else:
mask = None
h = makeHistogram(
self.hist_name,
columns,
self.axes,
data,
mask=mask,
)
return columns, [h]