Skip to content

Base

Helpers and reusable metric base classes for set-based comparisons.

Functions:

Name Description
_convert_dict_to_tuple

Convert dictionaries into hashable tuples while dropping ignored keys and None values.

Classes:

Name Description
SingleFieldMetric

Store the optional field name used by single-field metrics.

MetricWithPrepareEntryAsSet

Normalize metric inputs into comparable sets.

MetricWithTpFpFnEntries

Track tp/fp/fn entries per record for downstream metrics.

SingleFieldMetric(field=None, **kwargs)

Bases: Metric

Base class for metrics that operate on a single field of the prediction and reference.

Parameters:

Name Type Description Default
field str | None

Optional field name extracted from dictionary-like inputs.

None

Other Parameters:

Name Type Description
**kwargs

Additional keyword arguments forwarded to Metric.

Source code in src/kibad_llm/metrics/base.py
40
41
42
43
44
45
46
47
48
49
50
def __init__(self, field: str | None = None, **kwargs) -> None:
    """Store the optional field name used by the metric.

    Args:
        field: Optional field name extracted from dictionary-like inputs.

    Keyword Args:
        **kwargs: Additional keyword arguments forwarded to `Metric`.
    """
    super().__init__(**kwargs)
    self.field = field

MetricWithPrepareEntryAsSet(flatten_dicts=False, ignore_subfields=None, process_entry_func=None, process_entry_batch_func=None, **kwargs)

Bases: SingleFieldMetric

Base class for metrics that normalize entries into sets before comparison.

Attributes:

Name Type Description
field

Optional field name extracted from dictionary inputs before comparison.

flatten_dicts

Whether dictionary inputs are flattened before field extraction.

ignore_subfields

Subfield names ignored when converting dictionary values into hashable tuples.

Methods:

Name Description
_prepare_entry_as_set

Normalize one prediction or reference entry into a set.

Parameters:

Name Type Description Default
flatten_dicts bool

Whether to flatten nested dictionaries before further processing.

False
ignore_subfields dict[str, list] | None

Optional mapping from field names to subfield names that should be ignored when converting dictionaries into tuples.

None
process_entry_func Callable[[Hashable], Hashable] | None

Optional method to process each entry before normalization (e.g., for lowercasing or using an entity linking service). Has no effect if process_entry_batch_func is provided.

None
process_entry_batch_func Callable[[list[Hashable]], list[Hashable]] | None

Optional method to process a batch of entries before normalization. This is only effective for list entries. Takes precedence over process_entry_func.

None

Other Parameters:

Name Type Description
field

Optional field to extract from dictionary inputs.

Raises:

Type Description
ValueError

If field is not provided, but flatten_dicts is enabled (because the flattened dictionary has lists as values which cannot be converted to a set).

Source code in src/kibad_llm/metrics/base.py
 66
 67
 68
 69
 70
 71
 72
 73
 74
 75
 76
 77
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
def __init__(
    self,
    flatten_dicts: bool = False,
    ignore_subfields: dict[str, list] | None = None,
    process_entry_func: Callable[[Hashable], Hashable] | None = None,
    process_entry_batch_func: Callable[[list[Hashable]], list[Hashable]] | None = None,
    **kwargs,
) -> None:
    """Initialize the shared entry-normalization settings.

    Args:
        flatten_dicts: Whether to flatten nested dictionaries before further processing.
        ignore_subfields: Optional mapping from field names to subfield names that should be
            ignored when converting dictionaries into tuples.
        process_entry_func: Optional method to process each entry before normalization (e.g., for
            lowercasing or using an entity linking service). Has no effect if `process_entry_batch_func`
            is provided.
        process_entry_batch_func: Optional method to process a batch of entries before normalization. This
            is only effective for list entries. Takes precedence over process_entry_func.

    Keyword Args:
        field: Optional field to extract from dictionary inputs.

    Raises:
        ValueError: If `field` is not provided, but `flatten_dicts` is enabled (because the flattened
            dictionary has lists as values which cannot be converted to a set).

    """
    super().__init__(**kwargs)
    if process_entry_func is not None:
        if process_entry_batch_func is not None:
            logger.warning(
                "Both process_entry_func and process_entry_batch_func are provided. "
                "process_entry_batch_func will take precedence."
            )
        else:
            process_entry_batch_func = lambda batch: [process_entry_func(e) for e in batch]

    self.process_entry_batch_func = process_entry_batch_func
    self.ignore_subfields = []
    self.flatten_dicts = flatten_dicts
    if ignore_subfields is not None and self.field is not None:
        self.ignore_subfields = ignore_subfields.get(self.field, [])
    if self.field is None and self.flatten_dicts:
        raise ValueError(
            "flatten_dicts is enabled, but no field is specified. "
            "Please provide a field to extract from the flattened dictionary."
        )

MetricWithTpFpFnEntries(ignore_missing_entries=False, **kwargs)

Bases: MetricWithPrepareEntryAsSet

Base class for metrics that retain tp/fp/fn entries instead of only counts.

Attributes:

Name Type Description
ignore_missing_entries

Whether updates with an empty prediction or reference side should be skipped.

state

Mapping from tp, fp, and fn to sets of (record_id, entry) pairs.

Methods:

Name Description
reset

Clear the tracked entries and seen record ids.

state_count

Return tp/fp/fn counts derived from the tracked entries.

state_per_record

Group tracked entries by record id.

Parameters:

Name Type Description Default
ignore_missing_entries bool

If True, skip updates where either side normalizes to an empty set.

False

Other Parameters:

Name Type Description
field

Optional field to extract from dictionary inputs.

flatten_dicts

Whether to flatten nested dictionaries before further processing.

ignore_subfields

Optional mapping from field names to subfield names that should be ignored when converting dictionaries into tuples.

Source code in src/kibad_llm/metrics/base.py
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
def __init__(self, ignore_missing_entries: bool = False, **kwargs) -> None:
    """Initialize tp/fp/fn entry tracking.

    Args:
        ignore_missing_entries: If `True`, skip updates where either side normalizes to an
            empty set.

    Keyword Args:
        field: Optional field to extract from dictionary inputs.
        flatten_dicts: Whether to flatten nested dictionaries before further processing.
        ignore_subfields: Optional mapping from field names to subfield names that should be
            ignored when converting dictionaries into tuples.
    """
    super().__init__(**kwargs)
    self.ignore_missing_entries = ignore_missing_entries
    self.reset()

state_count property

Return tp/fp/fn counts derived from the current entry state.

Returns:

Type Description
dict[str, int]

A mapping with the keys tp, fp, and fn and the number of tracked entries

dict[str, int]

for each category.

state_per_record property

Group the current tp/fp/fn state by record id.

Returns:

Type Description
dict[Hashable, dict[str, set]]

A nested mapping from record id to a dictionary with the keys tp, fp, and fn

dict[Hashable, dict[str, set]]

and sets of entries for that record.

reset()

Reset the tracked tp/fp/fn entries and the set of seen record ids.

Source code in src/kibad_llm/metrics/base.py
202
203
204
205
def reset(self) -> None:
    """Reset the tracked tp/fp/fn entries and the set of seen record ids."""
    self.state: dict[str, set] = {"tp": set(), "fp": set(), "fn": set()}
    self._used_record_ids: set[Hashable] = set()