| 140 | } |
| 141 | |
| 142 | void initialize_sub_comm(int size, int rank, torch::Tensor& kvs_data, std::vector<int> ranks) |
| 143 | { |
| 144 | ccl::kvs::address_type main_addr; |
| 145 | if (rank != 0) { |
| 146 | memcpy(main_addr.data(), kvs_data.data_ptr(), main_addr.size()); |
| 147 | sub_kvs = ccl::create_kvs(main_addr); |
| 148 | } |
| 149 | _ccl_comms.push_back(ccl::create_communicator(size, rank, sub_kvs)); |
| 150 | group_to_comm_id[ranks] = _ccl_comms.size() - 1; |
| 151 | } |
| 152 | |
| 153 | ccl::datatype get_ccl_datatype(c10::ScalarType type) |
| 154 | { |