mindspore/docs/api/api_python/ops/mindspore.ops.Dilation2D.rst

59 lines
3.5 KiB
ReStructuredText
Raw Normal View History

2022-11-15 20:30:51 +08:00
mindspore.ops.Dilation2D
=========================
.. py:class:: mindspore.ops.Dilation2D(stride, dilation, pad_mode="SAME", data_format="NCHW")
计算4-D和3-D输入Tensor的灰度膨胀。
对输入的shape为 :math:`(N, C_{in}, H_{in}, W_{in})` 应用2-D膨胀其中
:math:`N` 为batch大小 :math:`H` 为高度, :math:`W`为宽度, :math:`C` 为通道数量。
给定kernel size :math:`ks = (h_{ker}, w_{ker})`, stride :math:`s = (s_0, s_1)`,和
dilation :math:`d = (d_0, d_1)` ,计算如下:
.. math::
\text{output}(N_i, C_j, h, w) = \max_{m=0, \ldots, h_{ker}-1} \max_{n=0, \ldots, w_{ker}-1}
\text{input}(N_i, C_j, s_0 \times h + d_0 \times m, s_1 \times w + d_1 \times n) + \text{filter}(C_j, m, n)
.. warning::
- 此算子为实验性算子。
- 如果输入数据类型为float32算子仍然按float16模式执行。
参数:
- **stride** (Union(inttuple[int])) - kernel移动的距离。
如果为一个int整数则表示了height和width共同的步长。
如果为两个int整数的元组则分别表示height和width的步长。
如果为四个int整数的元组则说明数据格式为 `NCHW` ,表示 `[1, 1, stride_height, stride_width]`
- **dilation** (Union(inttuple[int])) - 数据类型为int或者包含2个整数的元组或者包含4个整数的元组指定用于扩张卷积的膨胀速率。
如果设置为 :math:`k > 1` ,则每次抽样点跳过 :math:`k - 1` 个像素点。
其值必须大于等于1并且以输入的宽度和高度为边界。
- **pad_mode** (str可选) - 指定填充模式,可选模式有"same", "valid",默认值:"same"。大小写均支持。
- same采用完全方式。输出的宽度和高度和输入的一样。
- valid采用丢弃的方式。没有填充时候输出为最大的高度和宽度。额外的像素点将被丢弃。
- **data_format** (str可选) - 数据格式的值。目前只支持`NCHW`,默认值: `NCHW`
输入:
- **x** (Tensor) - 输入数据。一个四维Tensor数据类型为float16、float32shape必须为
:math:`(N, C_{in}, H_{in}, W_{in})`
- **filter** (Tensor) - 一个三维Tensor数据类型和输入 `x` 相同shape必须为
:math:`(C_{in}, H_{filter}, W_{filter})`
输出:
Tensor其值已经过dilation2D。shape为 :math:`(N, C_{out}, H_{out}, W_{out})`,未必和输入 `x` shape相同数据类型和输入 `x` 相同。
异常:
- **TypeError** - 如果输入 `x` 或者 `filter` 的数据类型不是uint8、uint16、uint32、uint64、int8、int16、
int32、int64、float16、float32、float64。
- **TypeError** - 如果参数 `stride` 或者 `dilation` 不是一个整数或者包含两个整数的元组或者包含四个整数的元组。
- **ValueError** - 如果参数 `stride` 或者 `dilation` 是一个元组并且它的长度不是2或者4。
- **ValueError** - 如果参数 `stride` 或者 `dilation` 是一个包含四个整数的元组它的shape不是 `(1, 1, height, width)`
- **ValueError** - 如果参数 `stride` 的取值范围不是`[1, 255]`
- **ValueError** - 如果参数 `dilation` 的值小于1。
- **ValueError** - 如果参数 `pad_mode` 不是 `same``valid``SAME` 或者 `VALID`
- **ValueError** - 如果参数 `data_format` 不是字符串`NCHW`