Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 16 additions & 3 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -129,8 +129,15 @@ if(USE_CUDA)
message(STATUS "Add USE_NCCL, use NCCL with CUDA")
list(APPEND CMAKE_MODULE_PATH ${PROJECT_SOURCE_DIR}/cmake)
find_package(NCCL REQUIRED)
add_compile_definitions(USE_NCCL=1)
target_link_libraries(infini_train_cuda_kernels PUBLIC nccl)
# add_compile_definitions(USE_NCCL=1)
# target_link_libraries(infini_train_cuda_kernels PUBLIC nccl)

# 不再用目录级 add_compile_definitions 影响所有target, 只标到真正需要NCCL的target
target_compile_definitions(infini_train_cuda_kernels PUBLIC USE_NCCL=1)
# 加头文件
target_include_directories(infini_train_cuda_kernels PUBLIC ${NCCL_INCLUDE_DIRS})
# 用FindNCCL找到完整库路径
target_link_libraries(infini_train_cuda_kernels PUBLIC ${NCCL_LIBRARIES})
endif()
endif()

Expand Down Expand Up @@ -168,7 +175,13 @@ if(USE_CUDA)
if(USE_NCCL)
# If your core library code also directly references NCCL symbols (not only kernels),
# keep this. Otherwise it's harmless.
target_link_libraries(infini_train PUBLIC nccl)

# target_link_libraries(infini_train PUBLIC nccl)

# nccl_impl.cc需要nccl.h和NCCL库
target_compile_definitions(infini_train PUBLIC USE_NCCL=1)
target_include_directories(infini_train PUBLIC ${NCCL_INCLUDE_DIRS})
target_link_libraries(infini_train PUBLIC ${NCCL_LIBRARIES})
endif()
endif()

Expand Down
10 changes: 9 additions & 1 deletion example/mnist/dataset.cc
Original file line number Diff line number Diff line change
Expand Up @@ -115,8 +115,16 @@ MNISTDataset::MNISTDataset(const std::string &dataset, bool train)
std::pair<std::shared_ptr<infini_train::Tensor>, std::shared_ptr<infini_train::Tensor>>
MNISTDataset::operator[](size_t idx) const {
CHECK_LT(idx, image_file_.dims[0]);
return {std::make_shared<infini_train::Tensor>(image_file_.tensor, idx * image_size_in_bytes_, image_dims_),
// image_file_.tensor 在构造函数里已被转换为 float32,
// 每个样本的字节数必须按当前 dtype 计算(784 * sizeof(float) = 3136),
// 不能用原始 uint8 文件的 784。
const size_t image_sample_bytes = image_file_.tensor.SizeInBytes() / image_file_.dims[0];

// return {std::make_shared<infini_train::Tensor>(image_file_.tensor, idx * image_size_in_bytes_, image_dims_),
// std::make_shared<infini_train::Tensor>(label_file_.tensor, idx * label_size_in_bytes_, label_dims_)};
return {std::make_shared<infini_train::Tensor>(image_file_.tensor, idx * image_sample_bytes, image_dims_),
std::make_shared<infini_train::Tensor>(label_file_.tensor, idx * label_size_in_bytes_, label_dims_)};

}

size_t MNISTDataset::Size() const { return image_file_.dims[0]; }
Loading