FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
cpp-taskflow/sandbox/parallel_dnn/tbb.cpp at master · WarfareCode/cpp-taskflow · GitHub
WarfareCode
cpp-taskflow
Repository navigation
Code
Pull requests
Actions
Projects
Security and quality
Insights
Expand file tree
Breadcrumbs
cpp-taskflow
/
sandbox
/
parallel_dnn
/
tbb.cpp
Copy path
More file actions
More file actions
Latest commit
History
History
History
104 lines (84 loc) · 2.64 KB
Breadcrumbs
cpp-taskflow
/
sandbox
/
parallel_dnn
/
tbb.cpp
Copy path
File metadata and controls
104 lines (84 loc) · 2.64 KB
Raw
Copy raw file
Download raw file
Open symbols panel
Edit and raw actions
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
#
include
"
dnn.hpp
"
#
include
<
tbb/task_scheduler_init.h
>
#
include
<
tbb/flow_graph.h
>
using
namespace
tbb
;
using
namespace
tbb
::flow
;
struct
TBB_DNNTrainingPattern
{
TBB_DNNTrainingPattern
() {
init_dnn
(dnn,
rand_rate
());
build_task_graph
();
}
void
train
() {
f_task->
try_put
(
continue_msg
());
G.
wait_for_all
();
}
void
build_task_graph
() {
f_task = std::make_unique<continue_node<continue_msg>>(G,
[&](
const
continue_msg&) {
forward_task
(dnn,
IMAGES
,
LABELS
);
}
);
for
(
int
j=dnn.
acts
.
size
()-
1
; j>=
0
; j--) {
//
backward propagation
auto
& b_task = backward_tasks.
emplace_back
(
std::make_unique<continue_node<continue_msg>>(G,
[&, i=j](
const
continue_msg&) {
backward_task
(dnn, i,
IMAGES
);
})
);
auto
& u_task = update_tasks.
emplace_back
(
std::make_unique<continue_node<continue_msg>>(G,
[&, i=j](
const
continue_msg&) {
dnn.
update
(i);
})
);
if
(j +
1u
== dnn.
acts
.
size
()) {
make_edge
(*f_task, *b_task);
}
else
{
make_edge
(*backward_tasks[backward_tasks.
size
()-
2
], *b_task);
}
make_edge
(*b_task, *u_task);
}
}
tbb::flow::graph G;
MNIST_DNN
dnn;
std::unique_ptr<continue_node<continue_msg>> f_task;
std::vector<std::unique_ptr<continue_node<continue_msg>>> backward_tasks;
std::vector<std::unique_ptr<continue_node<continue_msg>>> update_tasks;
};
void
run_tbb
(
const
unsigned
num_epochs,
const
unsigned
num_threads) {
tbb::task_scheduler_init
init
(num_threads);
auto
dnn_patterns = std::make_unique<TBB_DNNTrainingPattern[]>(
NUM_DNNS
);
auto
dnns = std::make_unique<std::unique_ptr<continue_node<continue_msg>>[]>(
NUM_DNNS
);
tbb::flow::graph parallel_dnn;
for
(
auto
i=
0u
; i<
NUM_DNNS
; i++) {
dnns[i] = std::make_unique<continue_node<continue_msg>>(parallel_dnn,
[&, id=i](
const
continue_msg&){
for
(
size_t
i=
0
; i<
NUM_ITERATIONS
; i++) {
dnn_patterns[id].
train
();
}
}
);
}
auto
sync_node = std::make_unique<continue_node<continue_msg>>(parallel_dnn,
[&](
const
continue_msg&) {
for
(
size_t
i=
0
; i<
NUM_DNNS
; i++) {
//
std::cout << "Validate " << i << "th NN: ";
dnn_patterns[i].
dnn
.
validate
(
TEST_IMAGES
,
TEST_LABELS
);
}
shuffle
(
IMAGES
,
LABELS
);
}
);
for
(
auto
i=
0u
; i<
NUM_DNNS
; i++) {
make_edge
(*(dnns[i]), *sync_node);
}
//
auto t1 = std::chrono::high_resolution_clock::now();
for
(
auto
i=
0u
; i<num_epochs; i++) {
for
(
auto
j=
0u
; j<
NUM_DNNS
; j++) {
dnns[j]->
try_put
(
continue_msg
());
}
parallel_dnn.
wait_for_all
();
//
report_runtime(t1);
}
}
Back
|
FazBrowse Home
|
New Git URL