| 49 | |
| 50 | template<typename T> |
| 51 | void CalcSumOfBlobs(KernelContext* ctx, const std::function<Blob*(const std::string&)>& BnInOp2Blob, |
| 52 | const PbRpf<std::string>& src_bns, const std::string& dst_bn) { |
| 53 | Blob* dst_blob = BnInOp2Blob(dst_bn); |
| 54 | std::unique_ptr<ep::primitive::Add> primitive = |
| 55 | ep::primitive::NewPrimitive<ep::primitive::AddFactory>(DeviceType::kCPU, |
| 56 | dst_blob->data_type()); |
| 57 | CHECK(primitive); |
| 58 | std::vector<const void*> srcs(src_bns.size()); |
| 59 | FOR_RANGE(size_t, i, 0, src_bns.size()) { |
| 60 | Blob* src_blob_i = BnInOp2Blob(src_bns.Get(i)); |
| 61 | srcs[i] = src_blob_i->dptr<T>(); |
| 62 | } |
| 63 | primitive->Launch(ctx->stream(), srcs.data(), srcs.size(), dst_blob->mut_dptr<T>(), |
| 64 | dst_blob->static_shape().elem_cnt()); |
| 65 | } |
| 66 | |
| 67 | void CopyFromFirstToOtherBlobs(KernelContext* ctx, |
| 68 | const std::function<Blob*(const std::string&)>& BnInOp2Blob, |