MCPcopy Create free account
hub / github.com/Oneflow-Inc/oneflow / Compute

Method Compute

oneflow/user/kernels/eager_ccl_kernel.cpp:382–489  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

380 private:
381 using user_op::OpKernel::Compute;
382 void Compute(user_op::KernelComputeContext* ctx, user_op::OpKernelState*,
383 const user_op::OpKernelCache* cache) const override {
384 auto* kernel_cache = dynamic_cast<const EagerCclS2SOpKernelCache*>(cache);
385 CHECK(kernel_cache != nullptr);
386 // NOTE(hanbinbin): Compute logic copy from _nccl_logical_s2s
387 const user_op::Tensor* in = ctx->Tensor4ArgNameAndIndex("in", 0);
388 user_op::Tensor* out = ctx->Tensor4ArgNameAndIndex("out", 0);
389 user_op::Tensor* tmp_buffer = ctx->Tensor4ArgNameAndIndex("tmp_buffer", 0);
390 const int64_t dtype_size = GetSizeOfDataType(in->data_type());
391 int64_t data_size = in->shape_view().elem_cnt() * dtype_size;
392 // NOTE: in (transpose)-> pack_to_ptr (all2all)-> unpack_from_ptr (transpose)-> out
393 const char* pack_to_ptr = in->dptr<char>();
394 char* unpack_from_ptr = out->mut_dptr<char>();
395 int64_t tmp_size = tmp_buffer->shape_view().elem_cnt();
396 CHECK_EQ(tmp_size, data_size * 2);
397
398 CHECK_EQ(in->data_type(), out->data_type());
399 const int64_t num_ranks = kernel_cache->parallel_desc()->parallel_num();
400 CHECK_EQ(in->shape_view().elem_cnt(), out->shape_view().elem_cnt())
401 << in->shape_view().ToString() << " vs " << out->shape_view().ToString();
402 const int64_t elem_cnt = in->shape_view().elem_cnt();
403 const int64_t in_split_axis = ctx->Attr<int64_t>("in_split_axis");
404 const int64_t out_split_axis = ctx->Attr<int64_t>("out_split_axis");
405
406 DimVector logical_shape_dim_vec;
407 in->shape_view().ToDimVector(&logical_shape_dim_vec);
408 logical_shape_dim_vec[in_split_axis] = logical_shape_dim_vec.at(in_split_axis) * num_ranks;
409
410 if (out_split_axis != 0) {
411 // Do pack. Need transpose in -> pack_to
412 // pack use temp buffer offset: [0, data_size]
413 pack_to_ptr = tmp_buffer->dptr<char>();
414 DimVector transpose_in_dim_vec = logical_shape_dim_vec;
415 CHECK_EQ(transpose_in_dim_vec.at(in_split_axis) % num_ranks, 0);
416 transpose_in_dim_vec[in_split_axis] = transpose_in_dim_vec.at(in_split_axis) / num_ranks;
417 CHECK_EQ(transpose_in_dim_vec.at(out_split_axis) % num_ranks, 0);
418 transpose_in_dim_vec[out_split_axis] = transpose_in_dim_vec.at(out_split_axis) / num_ranks;
419 transpose_in_dim_vec.insert(transpose_in_dim_vec.begin() + out_split_axis, num_ranks);
420 std::vector<int32_t> perm;
421 perm.emplace_back(out_split_axis);
422 FOR_RANGE(int64_t, i, 0, transpose_in_dim_vec.size()) {
423 if (i != out_split_axis) { perm.emplace_back(i); }
424 }
425 auto transpose = ep::primitive::NewPrimitive<ep::primitive::PermuteFactory>(
426 ctx->stream()->device_type(), transpose_in_dim_vec.size());
427 CHECK(transpose);
428 transpose->Launch(ctx->stream(), in->data_type(), transpose_in_dim_vec.size(),
429 transpose_in_dim_vec.data(), in->dptr(), perm.data(),
430 tmp_buffer->mut_dptr());
431 }
432
433 if (in_split_axis != 0) {
434 // Do unpack. Need transpose unpack_from -> out
435 // unpack use temp buffer offset: [tmp_size - data_size, tmp_size]
436 unpack_from_ptr = tmp_buffer->mut_dptr<char>() + (tmp_size - data_size);
437 }
438
439 {

Callers

nothing calls this directly

Calls 15

GetSizeOfDataTypeFunction · 0.85
ToDimVectorMethod · 0.80
insertMethod · 0.80
SendFunction · 0.70
RecvFunction · 0.70
data_typeMethod · 0.45
elem_cntMethod · 0.45
shape_viewMethod · 0.45
parallel_numMethod · 0.45
parallel_descMethod · 0.45

Tested by

no test coverage detected