FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
cpp-taskflow/benchmark/mnist/tbb.cpp at master · fcccode/cpp-taskflow · GitHub
fcccode
/
cpp-taskflow
Public
forked from
taskflow/taskflow
Notifications
You must be signed in to change notification settings
Fork
0
Star
0
Code
Pull requests
0
Actions
Projects
Security and quality
0
Insights
Additional navigation options
Code
Pull requests
Actions
Projects
Security and quality
Insights
Expand file tree
Breadcrumbs
cpp-taskflow
/
benchmark
/
mnist
/
tbb.cpp
Copy path
More file actions
More file actions
Latest commit
History
History
History
117 lines (98 loc) · 3.49 KB
Breadcrumbs
cpp-taskflow
/
benchmark
/
mnist
/
tbb.cpp
Copy path
File metadata and controls
117 lines (98 loc) · 3.49 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
105
106
107
108
109
110
111
112
113
114
115
#
include
"
dnn.hpp
"
#
include
<
memory
>
//
unique_ptr
#
include
<
tbb/task_scheduler_init.h
>
#
include
<
tbb/flow_graph.h
>
void
run_tbb
(
MNIST
& D,
unsigned
num_threads) {
using
namespace
tbb
;
using
namespace
tbb
::flow
;
tbb::task_scheduler_init
init
(num_threads);
tbb::flow::graph G;
std::vector<std::unique_ptr<continue_node<continue_msg>>> forward_tasks;
std::vector<std::unique_ptr<continue_node<continue_msg>>> backward_tasks;
std::vector<std::unique_ptr<continue_node<continue_msg>>> update_tasks;
std::vector<std::unique_ptr<continue_node<continue_msg>>> shuffle_tasks;
//
Number of parallel shuffle
const
auto
num_storage = num_threads;
const
auto
num_par_shf =
std::min
(num_storage, D.
epoch
);
std::vector<Eigen::MatrixXf>
mats
(num_par_shf, D.
images
);
std::vector<Eigen::VectorXi>
vecs
(num_par_shf, D.
labels
);
//
Create task flow graph
const
auto
iter_num = D.
images
.
rows
()/D.
batch_size
;
for
(
auto
e=
0u
; e<D.
epoch
; e++) {
for
(
auto
i=
0u
; i<iter_num; i++) {
auto
& f_task = forward_tasks.
emplace_back
(
std::make_unique<continue_node<continue_msg>>(
G,
[&, i=i, e=e%num_par_shf](
const
continue_msg&) {
forward_task
(D, i, e, mats, vecs);
}
)
);
if
(i !=
0
|| (i ==
0
&& e !=
0
)) {
auto
sz = update_tasks.
size
();
for
(
auto
j=
1u
; j<=D.
acts
.
size
() ;j++) {
make_edge
(*update_tasks[sz-j], *f_task);
}
}
for
(
int
j=D.
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, e=e%num_par_shf] (
const
continue_msg&) {
backward_task
(D, i, e, mats);
}
)
);
//
update weight
auto
& u_task = update_tasks.
emplace_back
(
std::make_unique<continue_node<continue_msg>>(
G,
[&, i=j] (
const
continue_msg&) {
D.
update
(i);
}
)
);
if
(j +
1u
== D.
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);
}
//
End of backward propagation
}
//
End of all iterations (task flow graph creation)
if
(e ==
0
) {
//
No need to shuffle in first epoch
shuffle_tasks.
emplace_back
(std::make_unique<continue_node<continue_msg>>(
G,
[](
const
continue_msg&){}
));
make_edge
(*shuffle_tasks.
back
(), *forward_tasks[forward_tasks.
size
()-iter_num]);
}
else
{
auto
& t = shuffle_tasks.
emplace_back
(
std::make_unique<continue_node<continue_msg>>(
G,
[&, e=e%num_par_shf](
const
continue_msg&) {
D.
shuffle
(mats[e], vecs[e], D.
images
.
rows
());
}
)
);
make_edge
(*t, *forward_tasks[forward_tasks.
size
()-iter_num]);
//
This shuffle task starts after belows finish
//
1. previous shuffle on the same storage
//
2. the last backward task of previous epoch on the same storage
if
(e >= num_par_shf) {
auto
prev_e = e - num_par_shf;
make_edge
(*shuffle_tasks[prev_e] ,*t);
int
task_id = (prev_e+
1
)*iter_num*D.
acts
.
size
() -
1
;
make_edge
(*backward_tasks[task_id], *t);
}
}
}
//
End of all epoch
for
(
size_t
i=
0
; i<num_par_shf; i++) {
shuffle_tasks[i]->
try_put
(
continue_msg
());
}
G.
wait_for_all
();
}
Back
|
FazBrowse Home
|
New Git URL