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

Method exec

dnn/src/cuda/svd/opr_impl.cpp:74–147  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

72}
73
74void SVDForwardImpl::exec(
75 _megdnn_tensor_in src, _megdnn_tensor_out u, _megdnn_tensor_out s,
76 _megdnn_tensor_out vt, _megdnn_workspace workspace) {
77 Param p = param();
78 check_exec(src.layout, u.layout, s.layout, vt.layout, workspace.size);
79
80 size_t block_cnt, m, n;
81 canonize_params(src.layout, &block_cnt, &m, &n);
82
83 auto wbundle = get_workspace_bundle(
84 block_cnt, m, n, src.layout.dtype.size(), workspace.raw_ptr);
85 auto handle = concrete_handle(this->handle());
86
87 bool need_transpose = m > n;
88 size_t min_mn = std::min(m, n);
89 size_t max_mn = std::max(m, n);
90 TensorND cur_u, cur_v;
91 signed char job = 'N'; // Do not compute singular vectors.
92 if (p.compute_uv) {
93 SmallVector<size_t> u_shape, vt_shape;
94 if (p.full_matrices) {
95 job = 'A'; // Compute all singular vectors.
96 u_shape = {block_cnt, m, m};
97 vt_shape = {block_cnt, n, n};
98 } else {
99 job = 'S'; // Compute first min(m, n) singular vectors.
100 u_shape = {block_cnt, m, min_mn};
101 vt_shape = {block_cnt, min_mn, n};
102 }
103 if (need_transpose) {
104 cur_u = {
105 wbundle.get_workspace(3).raw_ptr,
106 {transposed_shape(u_shape), dtype::Float32()}};
107 cur_v = {
108 wbundle.get_workspace(4).raw_ptr,
109 {transposed_shape(vt_shape), dtype::Float32()}};
110 } else {
111 cur_v = {u.raw_ptr(), u.layout.reshape(u_shape)};
112 cur_u = {vt.raw_ptr(), vt.layout.reshape(vt_shape)};
113 }
114 } else {
115 cur_u = cur_v = {nullptr, {{0, 0}, dtype::Float32()}};
116 }
117
118 TensorND inp_copy(
119 wbundle.get_workspace(0).raw_ptr,
120 {{block_cnt, min_mn, max_mn}, dtype::Float32()});
121 float* cusolver_ws = wbundle.get_workspace(1).ptr<float>();
122 size_t cusolver_ws_size = wbundle.get_workspace(1).size / sizeof(float);
123 int* info = wbundle.get_workspace(2).ptr<int>();
124 TensorND s_blk(s.raw_ptr(), s.layout.reshape({block_cnt, min_mn}));
125
126 if (need_transpose) {
127 ::transpose(handle, src, inp_copy);
128 } else {
129 handle->relayout_opr()->exec(src, inp_copy);
130 }
131

Callers 1

transposeFunction · 0.45

Calls 14

maxFunction · 0.85
transposed_shapeFunction · 0.85
cusolver_handleMethod · 0.80
transposeFunction · 0.70
paramFunction · 0.50
get_workspace_bundleFunction · 0.50
concrete_handleFunction · 0.50
minFunction · 0.50
sizeMethod · 0.45
handleMethod · 0.45
get_workspaceMethod · 0.45
raw_ptrMethod · 0.45

Tested by

no test coverage detected