Skip to content

Commit f67275a

Browse files
committed
refine operator/math/CMakeLists.txt, seperate im2col from math_function
1 parent 6720198 commit f67275a

File tree

2 files changed

+52
-41
lines changed

2 files changed

+52
-41
lines changed

paddle/fluid/operators/CMakeLists.txt

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -173,11 +173,11 @@ op_library(parallel_do_op DEPS executor)
173173
op_library(create_reader_op DEPS reader)
174174

175175
if (WITH_GPU)
176-
op_library(conv_op DEPS vol2col depthwise_conv)
176+
op_library(conv_op DEPS vol2col depthwise_conv im2col)
177177
else()
178-
op_library(conv_op DEPS vol2col)
178+
op_library(conv_op DEPS vol2col im2col)
179179
endif()
180-
op_library(conv_transpose_op DEPS vol2col)
180+
op_library(conv_transpose_op DEPS vol2col im2col)
181181

182182
# FIXME(typhoonzero): save/load depends lodtensor serialization functions
183183
op_library(save_op DEPS lod_tensor)
Lines changed: 49 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -1,46 +1,57 @@
11
add_subdirectory(detail)
22

3+
function(math_library TARGET)
4+
# math_library is a function to create math library.
5+
# The interface is the same as cc_library.
6+
# But it handle split GPU/CPU code and link some common library.
7+
set(cc_srcs)
8+
set(cu_srcs)
9+
set(math_common_deps device_context framework_proto)
10+
set(multiValueArgs SRCS DEPS)
11+
cmake_parse_arguments(math_library "${options}" "${oneValueArgs}"
12+
"${multiValueArgs}" ${ARGN})
13+
14+
if (EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/${TARGET}.cc)
15+
list(APPEND cc_srcs ${TARGET}.cc)
16+
endif()
17+
if (EXISTS ${CMAKE_CURRENT_SOURCE_DIR}/${TARGET}.cu)
18+
list(APPEND cu_srcs ${TARGET}.cu)
19+
endif()
20+
21+
if (WITH_GPU)
22+
nv_library(${TARGET} SRCS ${cc_srcs} ${cu_srcs} DEPS ${math_library_DEPS} ${math_common_deps})
23+
else()
24+
cc_library(${TARGET} SRCS ${cc_srcs} DEPS ${math_library_DEPS} ${math_common_deps})
25+
endif()
26+
endfunction()
27+
28+
math_library(math_function DEPS cblas)
29+
math_library(im2col)
30+
math_library(selected_rows_functor DEPS selected_rows)
31+
math_library(softmax)
32+
math_library(cross_entropy)
33+
math_library(pooling)
34+
math_library(sequence_pooling)
35+
math_library(vol2col)
36+
math_library(context_project)
37+
math_library(sequence2batch)
38+
math_library(sequence_padding)
39+
math_library(sequence_scale)
40+
math_library(maxouting)
41+
math_library(unpooling)
42+
math_library(cos_sim_functor)
43+
math_library(lstm_compute DEPS activation_functions)
44+
math_library(gru_compute DEPS activation_functions)
345
if(WITH_GPU)
4-
nv_library(math_function SRCS math_function.cc math_function.cu im2col.cc im2col.cu DEPS cblas device_context framework_proto)
5-
nv_test(math_function_gpu_test SRCS math_function_test.cu DEPS math_function tensor)
6-
nv_library(selected_rows_functor SRCS selected_rows_functor.cc selected_rows_functor.cu DEPS selected_rows math_function)
7-
nv_test(selected_rows_functor_gpu_test SRCS selected_rows_functor_test.cu DEPS selected_rows_functor)
8-
nv_library(softmax SRCS softmax.cc softmax.cu DEPS device_context)
9-
nv_library(cross_entropy SRCS cross_entropy.cc cross_entropy.cu DEPS device_context)
10-
nv_library(pooling SRCS pooling.cc pooling.cu DEPS device_context)
1146
nv_library(depthwise_conv SRCS depthwise_conv.cu DEPS device_context)
12-
nv_library(sequence_pooling SRCS sequence_pooling.cc sequence_pooling.cu DEPS device_context math_function)
13-
nv_library(vol2col SRCS vol2col.cc vol2col.cu DEPS device_context tensor)
14-
nv_library(context_project SRCS context_project.cc context_project.cu DEPS device_context math_function)
15-
nv_library(sequence2batch SRCS sequence2batch.cc sequence2batch.cu DEPS device_context tensor math_function)
16-
nv_library(sequence_padding SRCS sequence_padding.cc sequence_padding.cu DEPS lod_tensor device_context)
17-
nv_library(sequence_scale SRCS sequence_scale.cc sequence_scale.cu DEPS lod_tensor device_context)
18-
nv_library(lstm_compute SRCS lstm_compute.cc lstm_compute.cu DEPS device_context activation_functions)
19-
nv_library(maxouting SRCS maxouting.cc maxouting.cu DEPS device_context)
20-
nv_library(unpooling SRCS unpooling.cc unpooling.cu DEPS device_context)
21-
nv_library(gru_compute SRCS gru_compute.cc gru_compute.cu DEPS device_context activation_functions math_function)
22-
nv_library(cos_sim_functor SRCS cos_sim_functor.cc cos_sim_functor.cu DEPS device_context)
23-
else()
24-
cc_library(math_function SRCS math_function.cc im2col.cc DEPS cblas device_context framework_proto)
25-
cc_library(selected_rows_functor SRCS selected_rows_functor.cc DEPS selected_rows math_function)
26-
cc_library(softmax SRCS softmax.cc DEPS device_context)
27-
cc_library(cross_entropy SRCS cross_entropy.cc DEPS device_context)
28-
cc_library(pooling SRCS pooling.cc DEPS device_context)
29-
cc_library(sequence_pooling SRCS sequence_pooling.cc DEPS device_context math_function)
30-
cc_library(vol2col SRCS vol2col.cc DEPS device_context tensor)
31-
cc_library(context_project SRCS context_project.cc DEPS device_context math_function)
32-
cc_library(sequence2batch SRCS sequence2batch.cc DEPS device_context tensor math_function)
33-
cc_library(sequence_padding SRCS sequence_padding.cc DEPS lod_tensor device_context)
34-
cc_library(sequence_scale SRCS sequence_scale.cc DEPS lod_tensor device_context)
35-
cc_library(lstm_compute SRCS lstm_compute.cc DEPS device_context activation_functions)
36-
cc_library(maxouting SRCS maxouting.cc DEPS device_context)
37-
cc_library(unpooling SRCS unpooling.cc DEPS device_context)
38-
cc_library(gru_compute SRCS gru_compute.cc DEPS device_context activation_functions math_function)
39-
cc_library(cos_sim_functor SRCS cos_sim_functor.cc DEPS device_context)
4047
endif()
4148

42-
cc_test(math_function_test SRCS math_function_test.cc DEPS math_function tensor)
49+
cc_test(math_function_test SRCS math_function_test.cc)
4350
cc_test(selected_rows_functor_test SRCS selected_rows_functor_test.cc DEPS selected_rows_functor)
44-
cc_test(im2col_test SRCS im2col_test.cc DEPS math_function tensor)
45-
cc_test(vol2col_test SRCS vol2col_test.cc DEPS vol2col tensor)
51+
cc_test(im2col_test SRCS im2col_test.cc DEPS im2col)
52+
cc_test(vol2col_test SRCS vol2col_test.cc DEPS vol2col)
4653
cc_test(sequence_padding_test SRCS sequence_padding_test.cc DEPS sequence_padding)
54+
if(WITH_GPU)
55+
nv_test(math_function_gpu_test SRCS math_function_test.cu)
56+
nv_test(selected_rows_functor_gpu_test SRCS selected_rows_functor_test.cu DEPS selected_rows_functor)
57+
endif()

0 commit comments

Comments
 (0)