You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在Keras MirroredStrategy中控制有状态Metric的归约策略

How to Control Cross-GPU Reduction for Custom Stateful Keras Metrics

I totally get your frustration here—dealing with distributed training metrics can be tricky because the cross-GPU aggregation logic is hidden away from the Metric subclass itself, unlike Losses which have a reduction parameter. Let’s break down what’s happening and how you can fix it:

First, Clarify the Difference in Reduction Logic

You’re right that the reduction mentioned in Metric’s code is for aggregating values within a single GPU’s Metric instance (e.g., accumulating batch values over epochs). The cross-GPU aggregation you’re seeing happens at the distributed strategy level, where each GPU’s Metric result is summed together by default. This isn’t controlled by the Metric class directly—Keras handles this behind the scenes when using MirroredStrategy.

Solution: Manually Handle Cross-GPU Aggregation in Your Metric

The workaround is to take control of the cross-GPU reduction yourself in the result() method of your custom Metric. Here’s how to do it:

Example: Custom Mean Metric with Controlled Reduction

import tensorflow as tf
from tensorflow.keras.metrics import Metric

class CustomControlledMean(Metric):
    def __init__(self, name="custom_controlled_mean", **kwargs):
        super().__init__(name=name, **kwargs)
        # Track total sum and count of samples across batches (per GPU)
        self.total = self.add_weight(name="total", initializer="zeros")
        self.count = self.add_weight(name="count", initializer="zeros")
        # Get the active distributed strategy
        self.strategy = tf.distribute.get_strategy()

    def update_state(self, y_true, y_pred, sample_weight=None):
        # Calculate per-batch values (adjust this to your metric logic)
        batch_errors = tf.abs(y_true - y_pred)
        batch_sum = tf.reduce_sum(batch_errors)
        batch_count = tf.cast(tf.shape(y_pred)[0], tf.float32)

        # Apply sample weights if provided
        if sample_weight is not None:
            batch_sum = tf.reduce_sum(batch_errors * sample_weight)
            batch_count = tf.reduce_sum(sample_weight)

        # Update per-GPU accumulators
        self.total.assign_add(batch_sum)
        self.count.assign_add(batch_count)

    def result(self):
        # Step 1: Aggregate total and count across all GPUs using SUM
        aggregated_total = self.strategy.reduce(
            tf.distribute.ReduceOp.SUM, self.total, axis=None
        )
        aggregated_count = self.strategy.reduce(
            tf.distribute.ReduceOp.SUM, self.count, axis=None
        )
        # Step 2: Compute the final mean (or any reduction you want)
        return aggregated_total / aggregated_count

    def reset_state(self):
        # Reset per-GPU accumulators at start of epoch
        self.total.assign(0.0)
        self.count.assign(0.0)

Why This Works:

By explicitly calling self.strategy.reduce() with ReduceOp.SUM on both our accumulators, we get the global total sum and sample count across all GPUs. Then we can compute the mean (or any other operation like max, min, etc.) instead of relying on Keras’ default behavior of summing the individual GPU metric results.

Where to Find the Cross-GPU Reduction Logic in TensorFlow

If you want to dig into the source code:

  • The core aggregation logic for Keras metrics in distributed training lives in tensorflow/python/keras/engine/training/distributed_training_utils.py, specifically the _aggregate_metrics function. This function uses the strategy’s reduce method with ReduceOp.SUM by default to combine metric values from all replicas.
  • The MirroredStrategy’s reduce implementation is in tensorflow/python/distribute/mirrored_strategy.py within the reduce_to method, which handles synchronizing values across GPUs.

As you noted, there’s an existing discussion about this exact issue in the TensorFlow repo, where other developers have run into the same default sum behavior and used similar manual aggregation workarounds.

内容的提问来源于stack exchange,提问作者Tomasz Bartczak

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.29 10:47:28