MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / compute3d

Function compute3d

dnn/src/naive/convolution3d/helper.h:58–161  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

56
57template <typename stype, typename ftype, typename dtype, class Strategy>
58void compute3d(
59 _megdnn_tensor_in src, ftype* __restrict fptr, _megdnn_tensor_out dst,
60 const Convolution3D::CanonizedFilterMeta& filter_meta) {
61 size_t spatial_start, channel_pos;
62 using Format = param::Convolution3D::Format;
63 if (filter_meta.format == Format::NCDHW) {
64 spatial_start = 2;
65 channel_pos = 1;
66 } else {
67 megdnn_assert(filter_meta.format == Format::NDHWC, "invalid conv format");
68 spatial_start = 1;
69 channel_pos = 4;
70 }
71 auto N = src.layout.shape[0], ID = src.layout.shape[spatial_start],
72 IH = src.layout.shape[spatial_start + 1],
73 IW = src.layout.shape[spatial_start + 2];
74 auto FD = filter_meta.spatial[0], FH = filter_meta.spatial[1],
75 FW = filter_meta.spatial[2];
76 auto OC = dst.layout.shape[channel_pos], OD = dst.layout.shape[spatial_start],
77 OH = dst.layout.shape[spatial_start + 1],
78 OW = dst.layout.shape[spatial_start + 2];
79
80 size_t FS_G, FS_OC, FS_IC, FS_SPATIAL;
81 if (filter_meta.format == Format::NCDHW) {
82 // g, oc, ic, fd, fh, fw
83 FS_SPATIAL = 1;
84 FS_IC = FD * FH * FW;
85 FS_OC = FS_IC * filter_meta.icpg;
86 FS_G = FS_OC * filter_meta.ocpg;
87 } else {
88 // g, oc, fd, fh, fw, ic
89 megdnn_assert(filter_meta.format == Format::NDHWC, "invalid conv format");
90 FS_IC = 1;
91 FS_SPATIAL = filter_meta.icpg;
92 FS_OC = FS_SPATIAL * FD * FH * FW;
93 FS_G = FS_OC * filter_meta.ocpg;
94 }
95
96 int pd = filter_meta.padding[0], ph = filter_meta.padding[1],
97 pw = filter_meta.padding[2];
98 size_t sd = filter_meta.stride[0], sh = filter_meta.stride[1],
99 sw = filter_meta.stride[2];
100 int dd = filter_meta.dilation[0], dh = filter_meta.dilation[1],
101 dw = filter_meta.dilation[2];
102 stype* __restrict sptr = src.ptr<stype>();
103 dtype* __restrict dptr = dst.ptr<dtype>();
104
105 int d_offset = -pd, h_offset = -ph, w_offset = -pw;
106
107 if (filter_meta.should_flip) {
108 d_offset += filter_meta.dilated_spatial[0] - 1;
109 h_offset += filter_meta.dilated_spatial[1] - 1;
110 w_offset += filter_meta.dilated_spatial[2] - 1;
111 dd = -dd;
112 dh = -dh;
113 dw = -dw;
114 }
115

Callers

nothing calls this directly

Calls 2

onFunction · 0.85
nextMethod · 0.45

Tested by

no test coverage detected