@synalinks_export("synalinks.metrics.Accuracy")
class Accuracy(Metric):
"""Computes per-field token accuracy.
Formula (per field, Jaccard index over normalized tokens):
```python
accuracy = |y_true_tokens ∩ y_pred_tokens| / |y_true_tokens ∪ y_pred_tokens|
```
Its output range is `[0, 1]`. It operates at a word level
and can be used for **QA systems**.
If `y_true` and `y_pred` contain multiple fields the JSON object's
fields are flattened and the score computed for each one
independently before being averaged.
Example:
```python
program.compile(
metrics=[
synalinks.metrics.Accuracy(),
],
)
```
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
tokens 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 `y_true` tokens for each field).
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).
"""
direction = "up"
def __init__(
self,
average=None,
name="accuracy",
in_mask=None,
out_mask=None,
in_mask_pattern=None,
out_mask_pattern=None,
):
super().__init__(
name=name,
in_mask=in_mask,
out_mask=out_mask,
in_mask_pattern=in_mask_pattern,
out_mask_pattern=out_mask_pattern,
)
if average not in (None, "micro", "macro", "weighted"):
raise ValueError(
"Invalid `average` argument value. Expected one of: "
"[None, 'micro', 'macro', 'weighted']. "
f"Received: average={average}"
)
self.state = self.add_variable(
data_model=AccuracyState,
name="state_" + self.name,
)
self.average = average
self.axis = None
if self.average != "micro":
self.axis = 0
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: str(x), y_true.get_json()))
y_pred = tree.flatten(tree.map_structure(lambda x: str(x), y_pred.get_json()))
correct = []
total = []
intermediate_weights = []
# zip_longest, not zip: a prediction with fewer/more leaves than the
# gold must be scored against every leaf; plain zip would silently
# drop the unmatched ones and inflate the score. The "" fill has no
# tokens, so an unmatched leaf scores 0 over the other side's tokens.
for yt, yp in zip_longest(y_true, y_pred, fillvalue=""):
y_true_tokens = nlp_utils.normalize_and_tokenize(str(yt))
y_pred_tokens = nlp_utils.normalize_and_tokenize(str(yp))
common_tokens = set(y_true_tokens) & set(y_pred_tokens)
union_tokens = set(y_true_tokens) | set(y_pred_tokens)
correct.append(len(common_tokens))
total.append(len(union_tokens))
intermediate_weights.append(len(y_true_tokens))
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):
if self.state.get("correct") is None and self.state.get("total") is None:
return 0.0
correct = np.convert_to_tensor(self.state.get("correct"))
total = np.convert_to_tensor(self.state.get("total"))
# Keras/sklearn "micro": aggregate correct/total across all fields
# first, *then* compute the ratio. Without this collapse, "micro"
# would degenerate to a mean over per-field scores (i.e. macro).
if self.average == "micro":
correct = np.sum(correct)
total = np.sum(total)
score = np.divide(correct, np.add(total, backend.epsilon()))
return self._aggregate(score)
def _aggregate(self, score):
"""Apply `average` reduction over per-field scores."""
score = np.convert_to_tensor(score)
if self.average == "weighted":
intermediate_weights = self.state.get("intermediate_weights")
weights = np.divide(
intermediate_weights,
np.sum(intermediate_weights) + backend.epsilon(),
)
score = np.sum(score * weights)
elif self.average is not None: # [micro, macro]
score = np.mean(score, self.axis)
# numpy 1.25+ deprecates float() on a >0-D array even when size == 1,
# so go through .item() / .tolist() to always hand back Python scalars.
score_arr = np.convert_to_numpy(score)
if score_arr.size == 1:
return float(score_arr.item())
return [float(v) for v in score_arr.tolist()]
def get_config(self):
"""Return the serializable config of the metric.
Returns:
(dict): The config dict.
"""
config = {
"name": self.name,
"average": self.average,
}
base_config = super().get_config()
return {**base_config, **config}