| 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 | { |
nothing calls this directly
no test coverage detected