| 162 | |
| 163 | MGB_DYN_TYPE_OBJ_FINAL_IMPL(Host2DeviceCopy); |
| 164 | Host2DeviceCopy::Host2DeviceCopy( |
| 165 | ComputingGraph& graph, const std::shared_ptr<HostTensorND>& host_data, |
| 166 | const Param& param, const OperatorNodeConfig& config) |
| 167 | : Super{&graph, config, "h2d", {}}, m_param{param}, m_host_data{host_data} { |
| 168 | auto out_cn = m_host_data->comp_node(); |
| 169 | if (config.has_comp_node_set()) |
| 170 | out_cn = config.get_single_comp_node(); |
| 171 | mgb_assert(out_cn.valid(), "can not get output comp node"); |
| 172 | |
| 173 | if (param.allow_cpu_mem_fwd && |
| 174 | out_cn.mem_node() == CompNode::default_cpu().mem_node() && |
| 175 | host_data->comp_node().mem_node() == out_cn.mem_node()) { |
| 176 | m_fwd_host_mem = true; |
| 177 | dv_helper::add_output(*this, host_data->dtype()); |
| 178 | } else { |
| 179 | m_fwd_host_mem = false; |
| 180 | add_output(None)->dtype(host_data->dtype()); |
| 181 | } |
| 182 | add_equivalence_component<ScalarHash<void*>>(host_data.get()); |
| 183 | add_equivalence_component<PODHash<Param>>(&m_param); |
| 184 | |
| 185 | this->comp_node(out_cn); |
| 186 | |
| 187 | output(0)->add_flag(VarNode::Flag::ALLOW_EMPTY_SHAPE); |
| 188 | } |
| 189 | |
| 190 | const TensorShape& Host2DeviceCopy::get_output_shape() { |
| 191 | return m_host_data->shape(); |
nothing calls this directly
no test coverage detected