Skip to content

Commit b50bd46

Browse files
authored
vivado/vitis support sample broadcasting merge ops (fastmachinelearning#1426)
* vivado/vitis support sample broadcasting merge ops * fix
1 parent 71acc01 commit b50bd46

2 files changed

Lines changed: 26 additions & 15 deletions

File tree

hls4ml/backends/vivado/passes/merge_templates.py

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,8 @@
66

77
merge_config_template = """struct config{index} : nnet::merge_config {{
88
static const unsigned n_elem = {n_elem};
9+
static const unsigned n_elem1 = {n_elem1};
10+
static const unsigned n_elem2 = {n_elem2};
911
static const unsigned reuse_factor = {reuse};
1012
}};\n"""
1113

@@ -21,8 +23,12 @@ def __init__(self):
2123

2224
def format(self, node):
2325
params = self._default_config_params(node)
24-
params['n_elem'] = node.get_input_variable(node.inputs[0]).size_cpp()
25-
26+
params['n_elem1'] = node.get_input_variable(node.inputs[0]).size_cpp()
27+
params['n_elem2'] = node.get_input_variable(node.inputs[1]).size_cpp()
28+
params['n_elem'] = max(params['n_elem1'], params['n_elem2'])
29+
io_type = node.model.config.get_config_value('IOType')
30+
if io_type != 'io_parallel':
31+
assert params['n_elem1'] == params['n_elem2'], 'broadcasting merge not supported non-io_parallel'
2632
return self.template.format(**params)
2733

2834

hls4ml/templates/vivado/nnet_utils/nnet_merge.h

Lines changed: 18 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,8 @@
99
namespace nnet {
1010

1111
struct merge_config {
12-
static const unsigned n_elem = 10;
12+
static const unsigned n_elem1 = 10;
13+
static const unsigned n_elem2 = 10;
1314
static const unsigned reuse_factor = 1;
1415
};
1516

@@ -34,56 +35,60 @@ struct concat_config {
3435
};
3536

3637
template <class input1_T, class input2_T, class res_T, typename CONFIG_T>
37-
void add(input1_T data1[CONFIG_T::n_elem], input2_T data2[CONFIG_T::n_elem], res_T res[CONFIG_T::n_elem]) {
38+
void add(input1_T data1[CONFIG_T::n_elem1], input2_T data2[CONFIG_T::n_elem2], res_T res[CONFIG_T::n_elem]) {
3839
#pragma HLS PIPELINE
3940

4041
for (int ii = 0; ii < CONFIG_T::n_elem; ii++) {
41-
res[ii] = data1[ii] + data2[ii];
42+
res[ii] = data1[ii % CONFIG_T::n_elem1] + data2[ii % CONFIG_T::n_elem2];
4243
}
4344
}
4445

4546
template <class input1_T, class input2_T, class res_T, typename CONFIG_T>
46-
void subtract(input1_T data1[CONFIG_T::n_elem], input2_T data2[CONFIG_T::n_elem], res_T res[CONFIG_T::n_elem]) {
47+
void subtract(input1_T data1[CONFIG_T::n_elem1], input2_T data2[CONFIG_T::n_elem2], res_T res[CONFIG_T::n_elem]) {
4748
#pragma HLS PIPELINE
4849

4950
for (int ii = 0; ii < CONFIG_T::n_elem; ii++) {
50-
res[ii] = data1[ii] - data2[ii];
51+
res[ii] = data1[ii % CONFIG_T::n_elem1] - data2[ii % CONFIG_T::n_elem2];
5152
}
5253
}
5354

5455
template <class input1_T, class input2_T, class res_T, typename CONFIG_T>
55-
void multiply(input1_T data1[CONFIG_T::n_elem], input2_T data2[CONFIG_T::n_elem], res_T res[CONFIG_T::n_elem]) {
56+
void multiply(input1_T data1[CONFIG_T::n_elem1], input2_T data2[CONFIG_T::n_elem2], res_T res[CONFIG_T::n_elem]) {
5657
#pragma HLS PIPELINE
5758

5859
for (int ii = 0; ii < CONFIG_T::n_elem; ii++) {
59-
res[ii] = data1[ii] * data2[ii];
60+
res[ii] = data1[ii % CONFIG_T::n_elem1] * data2[ii % CONFIG_T::n_elem2];
6061
}
6162
}
6263

6364
template <class input1_T, class input2_T, class res_T, typename CONFIG_T>
64-
void average(input1_T data1[CONFIG_T::n_elem], input2_T data2[CONFIG_T::n_elem], res_T res[CONFIG_T::n_elem]) {
65+
void average(input1_T data1[CONFIG_T::n_elem1], input2_T data2[CONFIG_T::n_elem2], res_T res[CONFIG_T::n_elem]) {
6566
#pragma HLS PIPELINE
6667

6768
for (int ii = 0; ii < CONFIG_T::n_elem; ii++) {
68-
res[ii] = (data1[ii] + data2[ii]) * ap_ufixed<1, 0>(0.5);
69+
res[ii] = (data1[ii % CONFIG_T::n_elem1] + data2[ii % CONFIG_T::n_elem2]) * ap_ufixed<1, 0>(0.5);
6970
}
7071
}
7172

7273
template <class input1_T, class input2_T, class res_T, typename CONFIG_T>
73-
void maximum(input1_T data1[CONFIG_T::n_elem], input2_T data2[CONFIG_T::n_elem], res_T res[CONFIG_T::n_elem]) {
74+
void maximum(input1_T data1[CONFIG_T::n_elem1], input2_T data2[CONFIG_T::n_elem2], res_T res[CONFIG_T::n_elem]) {
7475
#pragma HLS PIPELINE
7576

7677
for (int ii = 0; ii < CONFIG_T::n_elem; ii++) {
77-
res[ii] = (data1[ii] > data2[ii]) ? static_cast<res_T>(data1[ii]) : static_cast<res_T>(data2[ii]);
78+
res[ii] = (data1[ii % CONFIG_T::n_elem1] > data2[ii % CONFIG_T::n_elem2])
79+
? static_cast<res_T>(data1[ii % CONFIG_T::n_elem1])
80+
: static_cast<res_T>(data2[ii % CONFIG_T::n_elem2]);
7881
}
7982
}
8083

8184
template <class input1_T, class input2_T, class res_T, typename CONFIG_T>
82-
void minimum(input1_T data1[CONFIG_T::n_elem], input2_T data2[CONFIG_T::n_elem], res_T res[CONFIG_T::n_elem]) {
85+
void minimum(input1_T data1[CONFIG_T::n_elem1], input2_T data2[CONFIG_T::n_elem2], res_T res[CONFIG_T::n_elem]) {
8386
#pragma HLS PIPELINE
8487

8588
for (int ii = 0; ii < CONFIG_T::n_elem; ii++) {
86-
res[ii] = (data1[ii] < data2[ii]) ? static_cast<res_T>(data1[ii]) : static_cast<res_T>(data2[ii]);
89+
res[ii] = (data1[ii % CONFIG_T::n_elem1] < data2[ii % CONFIG_T::n_elem2])
90+
? static_cast<res_T>(data1[ii % CONFIG_T::n_elem1])
91+
: static_cast<res_T>(data2[ii % CONFIG_T::n_elem2]);
8792
}
8893
}
8994

0 commit comments

Comments
 (0)