Handles calling metric functions. Arguments: outputs: List of outputs (predictions). targets: List of targets. skip_target_masks: Optional. List of boolean for whether the corresponding target should be ignored or not. sample_weights: Optional list of sample weig
(self,
outputs,
targets=None,
skip_target_masks=None,
sample_weights=None,
masks=None,
return_weighted_metrics=False,
return_weighted_and_unweighted_metrics=False)
| 2014 | return metric_results |
| 2015 | |
| 2016 | def _handle_metrics(self, |
| 2017 | outputs, |
| 2018 | targets=None, |
| 2019 | skip_target_masks=None, |
| 2020 | sample_weights=None, |
| 2021 | masks=None, |
| 2022 | return_weighted_metrics=False, |
| 2023 | return_weighted_and_unweighted_metrics=False): |
| 2024 | """Handles calling metric functions. |
| 2025 | |
| 2026 | Arguments: |
| 2027 | outputs: List of outputs (predictions). |
| 2028 | targets: List of targets. |
| 2029 | skip_target_masks: Optional. List of boolean for whether the corresponding |
| 2030 | target should be ignored or not. |
| 2031 | sample_weights: Optional list of sample weight arrays. |
| 2032 | masks: List of computed output mask values. |
| 2033 | return_weighted_metrics: Flag that indicates whether weighted metrics |
| 2034 | should be computed instead of unweighted metrics. This flag is ignored |
| 2035 | when `return_weighted_and_unweighted_metrics` is enabled. |
| 2036 | return_weighted_and_unweighted_metrics: Flag that is used to indicate |
| 2037 | whether both weighted and unweighted metrics should be computed. When |
| 2038 | this is not enabled, we use `return_weighted_metrics` param to indicate |
| 2039 | whether weighted or unweighted metrics should be returned. |
| 2040 | |
| 2041 | Returns: |
| 2042 | A list of metric result tensors. |
| 2043 | """ |
| 2044 | # TODO(scottzhu): Update this to use the new training_endpoints. Currently |
| 2045 | # the eager and graph logic is bit different. |
| 2046 | skip_target_masks = skip_target_masks or [False] * len(outputs) |
| 2047 | metric_results = [] |
| 2048 | with K.name_scope('metrics'): |
| 2049 | # Invoke all metrics added using `compile`. |
| 2050 | for i in range(len(outputs)): |
| 2051 | if skip_target_masks[i]: |
| 2052 | continue |
| 2053 | output = outputs[i] if outputs else None |
| 2054 | target = targets[i] if targets else None |
| 2055 | output_mask = masks[i] if masks else None |
| 2056 | |
| 2057 | if (return_weighted_and_unweighted_metrics or |
| 2058 | not return_weighted_metrics): |
| 2059 | metric_results.extend( |
| 2060 | self._handle_per_output_metrics(self._per_output_metrics[i], |
| 2061 | target, output, output_mask)) |
| 2062 | if return_weighted_and_unweighted_metrics or return_weighted_metrics: |
| 2063 | metric_results.extend( |
| 2064 | self._handle_per_output_metrics( |
| 2065 | self._per_output_weighted_metrics[i], |
| 2066 | target, |
| 2067 | output, |
| 2068 | output_mask, |
| 2069 | weights=sample_weights[i] if sample_weights else None)) |
| 2070 | return metric_results |
| 2071 | |
| 2072 | def _check_trainable_weights_consistency(self): |
| 2073 | """Check trainable weights count consistency. |
no test coverage detected