mindspore.ops.assign

mindspore.ops.assign(variable, value)[源代码]

为网络参数赋值。

variablevalue 遵循隐式类型转换规则,使数据类型一致。如果它们具有不同的数据类型,则低精度数据类型将转换为相对最高精度的数据类型。

参数:
  • variable (Parameter) - 网路参数。 \((N,*)\) ,其中 \(*\) 表示任意数量的附加维度,其秩应小于8。

  • value (Tensor) - 要分配的值,shape与 variable 相同。

返回:

Tensor,数据类型和shape与 variable 相同。

异常:
  • TypeError - 如果 variable 不是Parameter。

  • TypeError - 如果 value 不是Tensor。

  • RuntimeError - 如果 variablevalue 不支持参数的数据类型转换。

支持平台:

Ascend GPU CPU

样例:

>>> value = Tensor([2.0], mindspore.float32)
>>> variable = mindspore.Parameter(Tensor([1.0], mindspore.float32), name="variable")
>>> ops.assign(variable, value)
>>> print(variable.asnumpy())
[2.]