---
title: KerasDistributeOptimizer构造函数
description: "| 产品 | 是否支持 |"
url: https://www.hiascend.com/document/detail/zh/TensorFlowCommercial/latest/migration/tfmigr1/tfmigr1_tfadapi_0058.html
sourcePath: /source/zh/TensorFlowCommercial/900/migration/tfmigr1/tfmigr1_tfadapi_0058.html
indexId: 8cead68d04fe3e65a350f49796ed1bd90128e92ed269f20a048ec8c119e63b4b79
---
# KerasDistributeOptimizer构造函数

#### 产品支持情况

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


#### 功能说明

KerasDistributeOptimizer类的构造函数，用于包装用户使用tf.Keras构造的脚本中的单机训练优化器，构造NPU分布式训练优化器。


#### 函数原型

```
class KerasDistributeOptimizer(optimizer_v2.OptimizerV2):
    def __init__(self, optimizer, name="NpuKerasOptimizer", **kwargs)
```


#### 参数说明

| 参数名 | 输入/输出 | 描述 |
| --- | --- | --- |
| optimizer | 输入 | 用于梯度计算和更新权重的单机版训练优化器。 |
| name | 输入 | 优化器名称。 |


#### 返回值

返回KerasDistributeOptimizer类对象。


#### 调用示例

```
import tensorflow as tf
from npu_bridge.npu_init import *

model=xxx  
model.compile(loss='mean_squared_error', optimizer=KerasDistributeOptimizer(tf.keras.optimizers.SGD()))
```
