---
title: npu.distribute.npu_distributed_keras_optimizer_wrapper
description: "| 产品 | 是否支持 |"
url: https://www.hiascend.com/document/detail/zh/TensorFlowCommercial/latest/migration/tfmigr2/tfmigr2_tfadapi_0021.html
sourcePath: /source/zh/TensorFlowCommercial/900/migration/tfmigr2/tfmigr2_tfadapi_0021.html
indexId: 9963c1e9d66cf06166e6511d2e102b263607de42e49dbcc39640a4532514523779
---
# npu.distribute.npu_distributed_keras_optimizer_wrapper

#### 产品支持情况

| 产品 | 是否支持 |
| --- | --- |
| Atlas 350 加速卡 | √ |
| Atlas A3 训练系列产品 / Atlas A3 推理系列产品 | √ |
| Atlas A2 训练系列产品 / Atlas A2 推理系列产品 | √ |
| Atlas 200I/500 A2 推理产品 | ☓ |
| Atlas 推理系列产品 | ☓ |
| Atlas 训练系列产品 | √ |


#### 功能说明

在更新梯度前，添加NPU的allreduce操作对梯度进行聚合，然后再更新梯度。该接口仅在分布式场景下使用。


#### 函数原型

```
def npu_distributed_keras_optimizer_wrapper(optimizer, reduce_reduction="mean", fusion=1, fusion_id=-1, group="hccl_world_group")
```


#### 参数说明

| 参数名 | 输入/输出 | 描述 |
| --- | --- | --- |
| optimizer | 输入 | TensorFlow梯度训练优化器。 |
| reduce\_reduction | 输入 | String类型。 聚合运算的类型，可以为"mean","max","min","prod"或"sum"。 |
| fusion | 输入 | int类型。 allreduce算子融合标识，支持以下取值： 0：网络编译时，不会对该算子进行融合，即该allreduce算子不和其他allreduce算子融合。 1：网络编译时，对该算子按照梯度切分策略进行融合。 2：网络编译时，对allreduce算子按照相同的fusion\_id进行融合，即“fusion\_id”相同的allreduce算子之间会进行融合。 |
| fusion\_id | 输入 | int类型。 allreduce算子的融合id。 当“fusion”取值为“2”时，网络编译时会对相同fusion\_id的allreduce算子进行融合。 |
| group | 输入 | String类型，最大长度为128字节，含结束符。 group名称，可以为用户自定义group或者"hccl\_world\_group"。 |


#### 返回值

返回输入的optimizer。


#### 调用示例

```
import npu_device as npu
optimizer = tf.keras.optimizers.SGD()
optimizer = npu.distribute.npu_distributed_keras_optimizer_wrapper(optimizer) # 使用NPU分布式计算，更新梯度
model.compile(optimizer = optimizer)
```
