Skip to content

Utils

merge_references_into_predictions(predictions, references, allow_missing_references=False, allow_missing_predictions=False, verbose=False)

Create a new Dataset with entries "prediction" and "reference" by merging references into predictions based on matching IDs.

Parameters:

Name Type Description Default
predictions dict

Dataset containing prediction entries.

required
references dict

Dataset containing reference entries.

required
allow_missing_references bool

If True, allows predictions without corresponding references. This will fill missing references with empty dictionaries. If False, raises an error if any prediction is missing a reference.

False
allow_missing_predictions bool

If True, allows references without corresponding predictions. If False, raises an error if any reference is missing a prediction. IMPORTANT: In either case, evaluation is only performed if the prediction is present, so missing predictions will never be evaluated. However, support (=TP+FN) calculation will be affected.

False
verbose bool

If True, logs warnings for any missing references.

False

Returns: A new Dataset where each entry contains a "prediction" and its corresponding "reference".

Source code in src/kibad_llm/dataset/utils.py
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
def merge_references_into_predictions(
    predictions: dict,
    references: dict,
    allow_missing_references: bool = False,
    allow_missing_predictions: bool = False,
    verbose: bool = False,
) -> dict[Hashable, dict[str, dict]]:
    """Create a new Dataset with entries "prediction" and "reference" by merging references
    into predictions based on matching IDs.

    Args:
        predictions: Dataset containing prediction entries.
        references: Dataset containing reference entries.
        allow_missing_references: If True, allows predictions without corresponding references.
            This will fill missing references with empty dictionaries. If False, raises an error
            if any prediction is missing a reference.
        allow_missing_predictions: If True, allows references without corresponding predictions.
            If False, raises an error if any reference is missing a prediction.
            IMPORTANT: In either case, evaluation is only performed if the prediction is
            present, so missing predictions will never be evaluated. However, support (=TP+FN)
            calculation will be affected.
        verbose: If True, logs warnings for any missing references.
    Returns:
        A new Dataset where each entry contains a "prediction" and its corresponding "reference".
    """

    missing_references = set(predictions) - set(references)
    if missing_references:
        if not allow_missing_references:
            raise ValueError(f"Missing references for the following keys: {missing_references}")
        elif verbose:
            logger.warning(
                f"Missing references for the following keys: {missing_references}. "
                "Filling missing references with empty dictionaries."
            )
    missing_predictions = set(references) - set(predictions)
    if missing_predictions:
        if not allow_missing_predictions:
            raise ValueError(f"Missing predictions for the following keys: {missing_predictions}")
        elif verbose:
            logger.warning(
                f"Missing predictions for the following keys: {missing_predictions}. "
                "IMPORTANT: Evaluation is only performed if the prediction is present, "
                "so missing predictions will not be evaluated. This means that support (=TP+FN) "
                "is computed without the corresponding references."
            )

    merged_dataset = {
        k: {"prediction": predictions[k], "reference": references.get(k, {})} for k in predictions
    }

    if isinstance(predictions, DictWithMetadata):
        merged_dataset = DictWithMetadata(
            merged_dataset,
            metadata=predictions.metadata,
        )

    return merged_dataset