| 72 | } |
| 73 | |
| 74 | void 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 |
no test coverage detected