Skip to content

CategoricalAccuracy metric

Bases: Accuracy

Computes accuracy on list / categorical structures.

Formula (per field, Jaccard index over label sets):

accuracy = |y_true_labels  y_pred_labels| / |y_true_labels  y_pred_labels|

Its output range is [0, 1]. It operates at a label level and can be used for classification or retrieval pipelines.

Unlike Accuracy, this metric considers each element of the list (or the string value) as one label, comparing label sets rather than tokenized words.

If labels is provided, accumulation is performed per-label (sklearn-style): for each label L, a sample is "correct for L" when L's presence in y_true matches its presence in y_pred. This enables stable macro/weighted averaging across batches even when some labels are absent from a given sample, and lets result() return a {label: score} dict for average=None.

If labels is None, a single global set-Jaccard is computed over the pooled label values; in that mode average=None returns one scalar (use labels=... for a per-label breakdown).

Example:

# for single label classification

class ListClassification(synalinks.DataModel):
    label: Literal["label", "label_1", "label_2"]

# for multi label classification

class ListClassification(synalinks.DataModel):
    labels: List[Literal["label", "label_1", "label_2"]]

# or use it with retrieval pipelines, in that case make sure to mask
# the correct fields.

class AnswerWithReferences(synalinks.DataModel):
    sources: List[str]
    answer: str

Compilation example:

program.compile(
    metrics=[
        synalinks.metrics.CategoricalAccuracy(),
    ],
)

Parameters:

Name Type Description Default
average str

Type of averaging to be performed across per-field results in the multi-field case. Acceptable values are None, "micro", "macro" and "weighted". Defaults to None. If None, no averaging is performed and result() will return the score for each field. If "micro", compute the metric globally by aggregating label counts across all fields. If "macro", compute the metric for each field, and return their unweighted mean. If "weighted", compute the metric for each field, and return their mean weighted by support (the number of true labels per field).

None
labels list

(Optional) Explicit list of label names to track. When provided, accumulation is per-label across all batches and result() returns a {label: score} dict for average=None.

None
name str

(Optional) string name of the metric instance.

'categorical_accuracy'
in_mask list

(Optional) list of keys to keep to compute the metric.

None
out_mask list

(Optional) list of keys to remove to compute the metric.

None
in_mask_pattern str

(Optional) Regex pattern; fields whose names match are kept (combined with in_mask via OR).

None
out_mask_pattern str

(Optional) Regex pattern; fields whose names match are dropped (combined with out_mask via OR).

None
Source code in synalinks/src/metrics/accuracy_metrics.py
@synalinks_export("synalinks.metrics.CategoricalAccuracy")
class CategoricalAccuracy(Accuracy):
    """Computes accuracy on list / categorical structures.

    Formula (per field, Jaccard index over label sets):

    ```python
    accuracy = |y_true_labels ∩ y_pred_labels| / |y_true_labels ∪ y_pred_labels|
    ```

    Its output range is `[0, 1]`. It operates at a label level
    and can be used for **classification** or **retrieval pipelines**.

    Unlike `Accuracy`, this metric considers each element of the list
    (or the string value) as **one label**, comparing label sets rather than
    tokenized words.

    If `labels` is provided, accumulation is performed per-label (sklearn-style):
    for each label `L`, a sample is "correct for L" when L's presence in
    `y_true` matches its presence in `y_pred`. This enables stable
    `macro`/`weighted` averaging across batches even when some labels are
    absent from a given sample, and lets `result()` return a `{label: score}`
    dict for `average=None`.

    If `labels` is `None`, a single global set-Jaccard is computed over the
    pooled label values; in that mode `average=None` returns one scalar
    (use `labels=...` for a per-label breakdown).

    Example:

    ```python

    # for single label classification

    class ListClassification(synalinks.DataModel):
        label: Literal["label", "label_1", "label_2"]

    # for multi label classification

    class ListClassification(synalinks.DataModel):
        labels: List[Literal["label", "label_1", "label_2"]]

    # or use it with retrieval pipelines, in that case make sure to mask
    # the correct fields.

    class AnswerWithReferences(synalinks.DataModel):
        sources: List[str]
        answer: str

    ```


    Compilation example:

    ```python
    program.compile(
        metrics=[
            synalinks.metrics.CategoricalAccuracy(),
        ],
    )
    ```

    Args:
        average (str): Type of averaging to be performed across per-field results
            in the multi-field case.
            Acceptable values are `None`, `"micro"`, `"macro"` and
            `"weighted"`. Defaults to `None`.
            If `None`, no averaging is performed and `result()` will return
            the score for each field.
            If `"micro"`, compute the metric globally by aggregating
            label counts across all fields.
            If `"macro"`, compute the metric for each field, and return their
            unweighted mean.
            If `"weighted"`, compute the metric for each field, and return their
            mean weighted by support (the number of true labels per field).
        labels (list): (Optional) Explicit list of label names to track.
            When provided, accumulation is per-label across all batches and
            `result()` returns a `{label: score}` dict for `average=None`.
        name (str): (Optional) string name of the metric instance.
        in_mask (list): (Optional) list of keys to keep to compute the metric.
        out_mask (list): (Optional) list of keys to remove to compute the metric.
        in_mask_pattern (str): (Optional) Regex pattern; fields whose names match
            are kept (combined with ``in_mask`` via OR).
        out_mask_pattern (str): (Optional) Regex pattern; fields whose names match
            are dropped (combined with ``out_mask`` via OR).
    """

    def __init__(
        self,
        average=None,
        labels=None,
        name="categorical_accuracy",
        in_mask=None,
        out_mask=None,
        in_mask_pattern=None,
        out_mask_pattern=None,
    ):
        super().__init__(
            average=average,
            name=name,
            in_mask=in_mask,
            out_mask=out_mask,
            in_mask_pattern=in_mask_pattern,
            out_mask_pattern=out_mask_pattern,
        )
        if labels is not None:
            labels = [str(label) for label in labels]
        self.labels = labels

    async def update_state(self, y_true, y_pred):
        y_pred = tree.map_structure(lambda x: ops.convert_to_json_data_model(x), y_pred)
        y_true = tree.map_structure(lambda x: ops.convert_to_json_data_model(x), y_true)

        if self.in_mask or self.in_mask_pattern:
            y_pred = tree.map_structure(
                lambda x: (
                    x.in_mask(mask=self.in_mask, pattern=self.in_mask_pattern)
                    if x is not None
                    else x
                ),
                y_pred,
            )
            y_true = tree.map_structure(
                lambda x: (
                    x.in_mask(mask=self.in_mask, pattern=self.in_mask_pattern)
                    if x is not None
                    else x
                ),
                y_true,
            )
        if self.out_mask or self.out_mask_pattern:
            y_pred = tree.map_structure(
                lambda x: (
                    x.out_mask(mask=self.out_mask, pattern=self.out_mask_pattern)
                    if x is not None
                    else x
                ),
                y_pred,
            )
            y_true = tree.map_structure(
                lambda x: (
                    x.out_mask(mask=self.out_mask, pattern=self.out_mask_pattern)
                    if x is not None
                    else x
                ),
                y_true,
            )

        if y_true is None or y_pred is None:
            # A failed prediction yields `y_pred is None`; there is nothing to
            # compare, so skip the sample instead of calling `.get_json()` on None.
            return
        y_true = tree.flatten(tree.map_structure(lambda x: x, y_true.get_json()))
        y_pred = tree.flatten(tree.map_structure(lambda x: x, y_pred.get_json()))

        correct = []
        total = []
        intermediate_weights = []

        if self.labels is not None:
            y_true_set = {str(v) for v in y_true}
            y_pred_set = {str(v) for v in y_pred}
            for label in self.labels:
                t = label in y_true_set
                p = label in y_pred_set
                correct.append(1 if t == p else 0)
                total.append(1)
                intermediate_weights.append(1 if t else 0)
        else:
            # Set-Jaccard over the full pool of labels: position-independent,
            # so that `["a","b"]` vs `["b","a"]` scores 1.0. Produces a single
            # entry per call; per-label tracking requires `labels=...`.
            y_true_labels = [str(v) for v in y_true]
            y_pred_labels = [str(v) for v in y_pred]
            common_labels = set(y_true_labels) & set(y_pred_labels)
            union_labels = set(y_true_labels) | set(y_pred_labels)
            correct.append(len(common_labels))
            total.append(len(union_labels))
            intermediate_weights.append(len(y_true_labels))

        correct = np.convert_to_numpy(correct)
        total = np.convert_to_numpy(total)
        intermediate_weights = np.convert_to_numpy(intermediate_weights)

        current_correct = self.state.get("correct")
        if current_correct:
            correct = ragged_add(current_correct, correct)

        current_total = self.state.get("total")
        if current_total:
            total = ragged_add(current_total, total)

        current_intermediate_weights = self.state.get("intermediate_weights")
        if current_intermediate_weights:
            intermediate_weights = ragged_add(
                current_intermediate_weights, intermediate_weights
            )

        self.state.update(
            {
                "correct": correct.tolist(),
                "total": total.tolist(),
                "intermediate_weights": intermediate_weights.tolist(),
            }
        )

    def result(self):
        res = super().result()
        if self.labels is not None and self.average is None and isinstance(res, list):
            return {label: score for label, score in zip(self.labels, res)}
        return res

    def get_config(self):
        """Return the serializable config of the metric.

        Returns:
            (dict): The config dict.
        """
        config = {
            "labels": list(self.labels) if self.labels is not None else None,
            "name": self.name,
        }
        base_config = super().get_config()
        return {**base_config, **config}

get_config()

Return the serializable config of the metric.

Returns:

Type Description
dict

The config dict.

Source code in synalinks/src/metrics/accuracy_metrics.py
def get_config(self):
    """Return the serializable config of the metric.

    Returns:
        (dict): The config dict.
    """
    config = {
        "labels": list(self.labels) if self.labels is not None else None,
        "name": self.name,
    }
    base_config = super().get_config()
    return {**base_config, **config}