Skip to content

Commit 881c694

Browse files
authored
[PAL] Add SHMEM adaptor support (#522)
1 parent 02cd0f1 commit 881c694

72 files changed

Lines changed: 8012 additions & 2522 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

.github/workflows/unittest-adaptor.yml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -49,11 +49,11 @@ jobs:
4949
run: |
5050
cd /__w/FlagCX/FlagCX/test/unittest/adaptor
5151
export MPI_HOME=/usr/local/mpi
52-
make -j$(nproc)
52+
make -j$(nproc) USE_NVIDIA=1
5353
5454
- name: Run adaptor unit tests (requires GPU)
5555
run: |
5656
cd /__w/FlagCX/FlagCX/test/unittest/adaptor
5757
export MPI_HOME=/usr/local/mpi
5858
export LD_LIBRARY_PATH=/__w/FlagCX/FlagCX/build/lib:$LD_LIBRARY_PATH
59-
make run-unit
59+
make run-unit USE_NVIDIA=1

.github/workflows/unittest-core.yml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -49,11 +49,11 @@ jobs:
4949
run: |
5050
cd /__w/FlagCX/FlagCX/test/unittest/core
5151
export MPI_HOME=/usr/local/mpi
52-
make -j$(nproc)
52+
make -j$(nproc) USE_NVIDIA=1
5353
5454
- name: Run core unit tests
5555
run: |
5656
cd /__w/FlagCX/FlagCX/test/unittest/core
5757
export MPI_HOME=/usr/local/mpi
5858
export LD_LIBRARY_PATH=/__w/FlagCX/FlagCX/build/lib:$LD_LIBRARY_PATH
59-
make run-unit
59+
make run-unit USE_NVIDIA=1

.github/workflows/unittest-p2p.yml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -49,11 +49,11 @@ jobs:
4949
run: |
5050
cd /__w/FlagCX/FlagCX/test/unittest/p2p
5151
export MPI_HOME=/usr/local/mpi
52-
make -j$(nproc)
52+
make -j$(nproc) USE_NVIDIA=1
5353
5454
- name: Run p2p unit tests (requires IB hardware)
5555
run: |
5656
cd /__w/FlagCX/FlagCX/test/unittest/p2p
5757
export MPI_HOME=/usr/local/mpi
5858
export LD_LIBRARY_PATH=/__w/FlagCX/FlagCX/build/lib:$LD_LIBRARY_PATH
59-
make run-unit
59+
make run-unit USE_NVIDIA=1

.github/workflows/unittest-rma.yml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -49,11 +49,11 @@ jobs:
4949
run: |
5050
cd /__w/FlagCX/FlagCX/test/unittest/rma
5151
export MPI_HOME=/usr/local/mpi
52-
make -j$(nproc)
52+
make -j$(nproc) USE_NVIDIA=1
5353
5454
- name: Run RMA unit tests
5555
run: |
5656
cd /__w/FlagCX/FlagCX/test/unittest/rma
5757
export MPI_HOME=/usr/local/mpi
5858
export LD_LIBRARY_PATH=/__w/FlagCX/FlagCX/build/lib:$LD_LIBRARY_PATH
59-
make run-mpi
59+
make run-mpi USE_NVIDIA=1

.github/workflows/unittest-runner.yml

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -49,19 +49,19 @@ jobs:
4949
run: |
5050
cd /__w/FlagCX/FlagCX/test/unittest/runner
5151
export MPI_HOME=/usr/local/mpi
52-
make -j$(nproc)
52+
make -j$(nproc) USE_NVIDIA=1
5353
5454
- name: Run runner unit tests (no MPI)
5555
run: |
5656
cd /__w/FlagCX/FlagCX/test/unittest/runner
5757
export MPI_HOME=/usr/local/mpi
5858
export LD_LIBRARY_PATH=/__w/FlagCX/FlagCX/build/lib:$LD_LIBRARY_PATH
59-
make run-unit
59+
make run-unit USE_NVIDIA=1
6060
6161
- name: Run runner MPI collective tests
6262
run: |
6363
cd /__w/FlagCX/FlagCX/test/unittest/runner
6464
export MPI_HOME=/usr/local/mpi
6565
export PATH=$MPI_HOME/bin:$PATH
6666
export LD_LIBRARY_PATH=/__w/FlagCX/FlagCX/build/lib:$LD_LIBRARY_PATH
67-
make run-mpi
67+
make run-mpi USE_NVIDIA=1

.github/workflows/unittest-service.yml

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -49,11 +49,11 @@ jobs:
4949
run: |
5050
cd /__w/FlagCX/FlagCX/test/unittest/service
5151
export MPI_HOME=/usr/local/mpi
52-
make -j$(nproc)
52+
make -j$(nproc) USE_NVIDIA=1
5353
5454
- name: Run service unit tests
5555
run: |
5656
cd /__w/FlagCX/FlagCX/test/unittest/service
5757
export MPI_HOME=/usr/local/mpi
5858
export LD_LIBRARY_PATH=/__w/FlagCX/FlagCX/build/lib:$LD_LIBRARY_PATH
59-
make run-unit
59+
make run-unit USE_NVIDIA=1

.github/workflows/unittest-symmem.yml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -49,7 +49,7 @@ jobs:
4949
run: |
5050
cd /__w/FlagCX/FlagCX/test/unittest/symmem
5151
export MPI_HOME=/usr/local/mpi
52-
make -j$(nproc)
52+
make -j$(nproc) USE_NVIDIA=1
5353
5454
- name: Run symmem tests
5555
run: |

Makefile

Lines changed: 32 additions & 133 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,10 @@ USE_SUNRISE ?= 0
2424
USE_PPU ?= 0
2525
COMPILE_KERNEL ?= 0
2626

27+
# Device API backend selection
28+
USE_SHMEM ?= 0
29+
SHMEM_HOME ?= /usr/local/nvshmem
30+
2731
# set to empty if not provided
2832
DEVICE_HOME ?=
2933
CCL_HOME ?=
@@ -141,142 +145,34 @@ COMPILE_KERNEL_HOST_FLAG=
141145
COMPILE_KERNEL_FLAG =
142146
HOST_COMPILER ?= g++
143147
ifeq ($(USE_NVIDIA), 1)
144-
include makefiles/nvidia_gencode.mk
145-
DEVICE_LIB = $(DEVICE_HOME)/lib64
146-
DEVICE_INCLUDE = $(DEVICE_HOME)/include $(DEVICE_HOME)/include/cccl
147-
DEVICE_LINK = -lcudart -lcuda
148-
DEVICE_PLATFORM = CUDA
149-
DEVICE_COMPILER = $(DEVICE_HOME)/bin/nvcc
150-
DEVICE_COMPILE_FLAG = -c --cudart=shared -Xcompiler -fPIC -MMD -MP -rdc=true -g $(DEVICE_COMPILER_GENCODE)
151-
DEVICE_LINK_FLAG = --cudart=shared -Xcompiler -fPIC $(DEVICE_COMPILER_GENCODE)
152-
DEVICE_FILE_EXTENSION = cu
153-
CCL_LIB = $(CCL_HOME)/lib
154-
CCL_INCLUDE = $(CCL_HOME)/include
155-
CCL_LINK = -lnccl
156-
ADAPTOR_FLAG = -DUSE_NVIDIA_ADAPTOR
157-
ifeq ($(NVCC_GENCODE_MULTICAST_UNSUPPORTED), 1)
158-
ADAPTOR_FLAG += -DNVCC_GENCODE_MULTICAST_UNSUPPORTED
159-
endif
148+
include makefiles/nvidia.mk
160149
else ifeq ($(USE_ASCEND), 1)
161-
DEVICE_LIB = $(DEVICE_HOME)/lib64
162-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
163-
DEVICE_LINK = -lascendcl
164-
CCL_LIB = $(CCL_HOME)/lib64
165-
CCL_INCLUDE = $(CCL_HOME)/include
166-
CCL_LINK = -lhccl
167-
ADAPTOR_FLAG = -DUSE_ASCEND_ADAPTOR
150+
include makefiles/ascend.mk
168151
else ifeq ($(USE_ILUVATAR_COREX), 1)
169-
DEVICE_LIB = $(DEVICE_HOME)/lib
170-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
171-
DEVICE_LINK = -lcudart -lcuda
172-
CCL_LIB = $(CCL_HOME)/lib
173-
CCL_INCLUDE = $(CCL_HOME)/include
174-
CCL_LINK = -lnccl
175-
ADAPTOR_FLAG = -DUSE_ILUVATAR_COREX_ADAPTOR
152+
include makefiles/iluvatar_corex.mk
176153
else ifeq ($(USE_CAMBRICON), 1)
177-
DEVICE_LIB = $(DEVICE_HOME)/lib64
178-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
179-
DEVICE_LINK = -lcnrt
180-
CCL_LIB = $(CCL_HOME)/lib64
181-
CCL_INCLUDE = $(CCL_HOME)/include
182-
CCL_LINK = -lcncl
183-
ADAPTOR_FLAG = -DUSE_CAMBRICON_ADAPTOR
154+
include makefiles/cambricon.mk
184155
else ifeq ($(USE_METAX), 1)
185-
DEVICE_LIB = $(DEVICE_HOME)/lib64
186-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
187-
CCL_LIB = $(CCL_HOME)/lib64
188-
CCL_INCLUDE = $(CCL_HOME)/include
189-
CCL_LINK = -lmccl
190-
ADAPTOR_FLAG = -DUSE_METAX_ADAPTOR
156+
include makefiles/metax.mk
191157
else ifeq ($(USE_MUSA), 1)
192-
DEVICE_LIB = $(DEVICE_HOME)/lib
193-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
194-
CCL_LIB = $(CCL_HOME)/lib
195-
CCL_INCLUDE = $(CCL_HOME)/include
196-
CCL_LINK = -lmccl -lmusa
197-
ADAPTOR_FLAG = -DUSE_MUSA_ADAPTOR
158+
include makefiles/musa.mk
198159
else ifeq ($(USE_KUNLUNXIN), 1)
199-
DEVICE_LIB = $(DEVICE_HOME)/so
200-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
201-
DEVICE_LINK = -lxpurt -lcudart
202-
CCL_LIB = $(CCL_HOME)/so
203-
CCL_INCLUDE = $(CCL_HOME)/include
204-
CCL_LINK = -lbkcl
205-
ADAPTOR_FLAG = -DUSE_KUNLUNXIN_ADAPTOR
160+
include makefiles/kunlunxin.mk
206161
else ifeq ($(USE_DU), 1)
207-
DEVICE_LIB = $(DEVICE_HOME)/lib64
208-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
209-
DEVICE_LINK = -lcudart -lcuda
210-
CCL_LIB = $(CCL_HOME)/lib64
211-
CCL_INCLUDE = $(CCL_HOME)/include
212-
CCL_LINK = -lnccl
213-
ADAPTOR_FLAG = -DUSE_DU_ADAPTOR
214-
DEVICE_PLATFORM = DU
215-
DEVICE_COMPILER = $(DEVICE_HOME)/bin/nvcc
216-
DEVICE_COMPILE_FLAG = -c --cudart=shared -Xcompiler -fPIC -MMD -MP -rdc=true -g
217-
DEVICE_LINK_FLAG = --cudart=shared -Xcompiler -fPIC
218-
DEVICE_FILE_EXTENSION = cu
162+
include makefiles/du.mk
219163
else ifeq ($(USE_AMD), 1)
220-
DEVICE_LIB = $(DEVICE_HOME)/lib
221-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
222-
DEVICE_LINK = -lhiprtc
223-
CCL_LIB = $(CCL_HOME)/lib
224-
CCL_INCLUDE = $(CCL_HOME)/include/rccl
225-
CCL_LINK = -lrccl
226-
ADAPTOR_FLAG = -DUSE_AMD_ADAPTOR -D__HIP_PLATFORM_AMD__
164+
include makefiles/amd.mk
227165
else ifeq ($(USE_TSM), 1)
228-
DEVICE_LIB = $(DEVICE_HOME)/lib
229-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
230-
DEVICE_LINK = -lhpgr
231-
CCL_LIB = $(CCL_HOME)/lib
232-
CCL_INCLUDE = $(CCL_HOME)/include
233-
CCL_LINK = -ltccl
234-
ADAPTOR_FLAG = -DUSE_TSM_ADAPTOR
166+
include makefiles/tsm.mk
235167
else ifeq ($(USE_ENFLAME), 1)
236-
DEVICE_LIB = $(DEVICE_HOME)/lib
237-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
238-
DEVICE_LINK = -ltopsrt
239-
CCL_LIB = $(CCL_HOME)/lib
240-
CCL_INCLUDE = $(CCL_HOME)/include
241-
CCL_LINK = -leccl
242-
ADAPTOR_FLAG = -DUSE_ENFLAME_ADAPTOR
168+
include makefiles/enflame.mk
243169
else ifeq ($(USE_SUNRISE), 1)
244-
DEVICE_LIB = $(DEVICE_HOME)/targets/linux-x86_64/lib
245-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
246-
DEVICE_LINK = -ltangrt_shared
247-
CCL_LIB = $(CCL_HOME)/lib/linux-x86_64
248-
CCL_INCLUDE = $(CCL_HOME)/include
249-
CCL_LINK = -lpccl
250-
ADAPTOR_FLAG = -DUSE_SUNRISE_ADAPTOR
170+
include makefiles/sunrise.mk
251171
else ifeq ($(USE_PPU), 1)
252-
DEVICE_LIB = $(DEVICE_HOME)/lib64
253-
DEVICE_INCLUDE = $(DEVICE_HOME)/include
254-
DEVICE_LINK = -lcudart -lcuda
255-
CCL_LIB = $(CCL_HOME)/lib64
256-
CCL_INCLUDE = $(CCL_HOME)/include
257-
CCL_LINK = -lnccl
258-
ADAPTOR_FLAG = -DUSE_PPU_ADAPTOR
172+
include makefiles/ppu.mk
259173
else
260-
DEVICE_LIB = $(DEVICE_HOME)/lib64
261-
DEVICE_INCLUDE = $(DEVICE_HOME)/include $(DEVICE_HOME)/include/cccl
262-
DEVICE_LINK = -lcudart -lcuda
263-
DEVICE_PLATFORM = CUDA
264-
DEVICE_COMPILER = $(DEVICE_HOME)/bin/nvcc
265-
DEVICE_COMPILE_FLAG = -c --cudart=shared -Xcompiler -fPIC -MMD -MP -rdc=true -g $(DEVICE_COMPILER_GENCODE)
266-
DEVICE_LINK_FLAG = --cudart=shared -Xcompiler -fPIC $(DEVICE_COMPILER_GENCODE)
267-
DEVICE_FILE_EXTENSION = cu
268-
CCL_LIB = $(CCL_HOME)/lib
269-
CCL_INCLUDE = $(CCL_HOME)/include
270-
CCL_LINK = -lnccl
271-
ADAPTOR_FLAG = -DUSE_NVIDIA_ADAPTOR
272-
ifeq ($(NVCC_GENCODE_MULTICAST_UNSUPPORTED), 1)
273-
ADAPTOR_FLAG += -DNVCC_GENCODE_MULTICAST_UNSUPPORTED
274-
endif
275-
USE_NVIDIA := 1
276-
endif
277-
278-
ifeq ($(FORCE_DEFAULT_PATH), 1)
279-
ADAPTOR_FLAG += -DFLAGCX_FORCE_DEFAULT_PATH
174+
include makefiles/nvidia.mk
175+
USE_NVIDIA := 1
280176
endif
281177

282178
ifeq ($(USE_GLOO), 1)
@@ -341,11 +237,18 @@ BUILD_PUBLIC_HEADERS := $(PUBLIC_HEADERS:flagcx/include/%=$(BUILD_INCDIR)/%)
341237
INCLUDEDIR := \
342238
$(abspath flagcx/include) \
343239
$(abspath flagcx/adaptor/include) \
240+
$(abspath flagcx/adaptor/device_api) \
241+
$(abspath flagcx/adaptor/shmem) \
344242
$(abspath flagcx/runner/include) \
345243
$(abspath flagcx/core/include) \
346244
$(abspath flagcx/service/include) \
347245
$(abspath third-party/json/single_include)
348246

247+
# Append NVSHMEM include path (must come after INCLUDEDIR := assignment)
248+
ifeq ($(USE_SHMEM), 1)
249+
INCLUDEDIR += $(SHMEM_HOME)/include
250+
endif
251+
349252
LIBSRCFILES:= \
350253
$(wildcard flagcx/*.cc) \
351254
$(wildcard flagcx/adaptor/*.cc) \
@@ -357,18 +260,14 @@ LIBSRCFILES:= \
357260
$(wildcard flagcx/core/*.cc) \
358261
$(wildcard flagcx/service/*.cc)
359262

263+
# Platform .mk provides extra sources (device_api backend, shmem adaptor)
264+
LIBSRCFILES += $(PLATFORM_EXTRA_SRCS)
265+
360266
ifeq ($(COMPILE_KERNEL), 1)
361-
DEVSRCFILES:= \
362-
$(wildcard flagcx/kernels/*.$(DEVICE_FILE_EXTENSION))
363-
ifneq ($(USE_NVIDIA), 1)
364-
EXCLUDE_SOURCES := custom_allreduce.cu
365-
else
366-
EXCLUDE_SOURCES :=
367-
endif
368-
DEVSRCFILES := $(filter-out flagcx/kernels/$(EXCLUDE_SOURCES), $(DEVSRCFILES))
369-
DEVOBJ:= $(DEVSRCFILES:%.$(DEVICE_FILE_EXTENSION)=$(OBJDIR)/%.o)
267+
DEVSRCFILES := $(PLATFORM_KERNEL_SRCS)
268+
DEVOBJ := $(DEVSRCFILES:%.$(DEVICE_FILE_EXTENSION)=$(OBJDIR)/%.o)
370269
endif
371-
LIBOBJ:= $(LIBSRCFILES:%.cc=$(OBJDIR)/%.o)
270+
LIBOBJ := $(LIBSRCFILES:%.cc=$(OBJDIR)/%.o)
372271

373272
TARGET = libflagcx.so
374273
all: $(LIBDIR)/$(TARGET) $(BUILD_PUBLIC_HEADERS)

0 commit comments

Comments
 (0)