• Home
  • Features
  • Pricing
  • Docs
  • Announcements
  • Sign In

daisytuner / docc / 26831349290

02 Jun 2026 03:50PM UTC coverage: 61.29% (-0.01%) from 61.302%
26831349290

Pull #725

github

web-flow
Merge a7e2175c0 into 887730e20
Pull Request #725: Tensor node backport

932 of 1642 new or added lines in 52 files covered. (56.76%)

92 existing lines in 33 files now uncovered.

35584 of 58058 relevant lines covered (61.29%)

11020.18 hits per line

Source File
Press 'n' to go to next uncovered line, 'b' for previous

95.16
/opt/src/targets/cuda/math/tensor/conv_expander.cpp
1
#include "sdfg/targets/cuda/math/tensor/conv_expander.h"
2
#include "sdfg/analysis/scope_analysis.h"
3
#include "sdfg/data_flow/access_node.h"
4
#include "sdfg/data_flow/library_nodes/math/blas/gemm_node.h"
5
#include "sdfg/data_flow/library_nodes/stdlib/free.h"
6
#include "sdfg/data_flow/library_nodes/stdlib/malloc.h"
7
#include "sdfg/structured_control_flow/block.h"
8
#include "sdfg/structured_control_flow/sequence.h"
9
#include "sdfg/types/pointer.h"
10
#include "sdfg/types/scalar.h"
11
#include "sdfg/types/tensor.h"
12

13
namespace sdfg {
14
namespace offloading {
15

16
bool CudaConvExpander::expand(builder::StructuredSDFGBuilder& builder, analysis::AnalysisManager& analysis_manager) {
4✔
17
    return expand_conv(builder, analysis_manager, node_);
4✔
18
}
4✔
19

20
bool CudaConvExpander::expand_conv(
21
    builder::StructuredSDFGBuilder& builder, analysis::AnalysisManager& analysis_manager, math::tensor::ConvNode& node
22
) {
4✔
23
    auto& dfg = node.get_parent();
4✔
24
    math::tensor::ConvNode::ConvExpandPrerequisits b;
4✔
25
    if (!node.check_expandable(dfg, analysis_manager, b)) {
4✔
26
        return false;
1✔
27
    }
1✔
28

29
    if (!symbolic::eq(node.group(), symbolic::one())) {
3✔
30
        // If there are no groups (i.e., group == 1), then we can do im2row with one GEMM.
31
        // Else, no CUDA specific expand is needed
32
        return false;
1✔
33
    }
1✔
34

35

36
    types::Scalar base_type(node.primitive_type(dfg));
2✔
37
    math::blas::BLAS_Precision precision = node.get_blas_precision(base_type);
2✔
38

39
    // Create new sequence for expansion
40
    auto& new_sequence = builder.add_sequence_before(
2✔
41
        *b.block_parent, *b.block, b.block_parent->at(b.block_index).second.assignments(), b.block->debug_info()
2✔
42
    );
2✔
43

44
    // Dimensions, i.e., 1D, 2D, 3D, ...
45
    size_t dims = node.kernel_shape().size();
2✔
46
    symbolic::MultiExpression out_shape = node.get_out_shape();
2✔
47
    types::Scalar indvar_type(types::PrimitiveType::Int64);
2✔
48

49

50
    /* ===== No groups ====================================================================== */
51

52
    // Add patches container with malloc
53
    symbolic::Expression patches_size = symbolic::mul(node.shape()[0], node.shape()[1]);
2✔
54
    for (size_t i = 0; i < dims; i++) {
5✔
55
        patches_size = symbolic::mul(patches_size, symbolic::mul(node.kernel_shape()[i], out_shape[i]));
3✔
56
    }
3✔
57
    types::Pointer patches_type(base_type);
2✔
58
    auto patches_container = builder.find_new_name("_patches");
2✔
59
    builder.add_container(patches_container, patches_type);
2✔
60
    auto [patches_malloc_block, patches_malloc_node] = stdlib::add_malloc_block(
2✔
61
        builder,
2✔
62
        new_sequence,
2✔
63
        patches_container,
2✔
64
        symbolic::mul(patches_size, symbolic::size_of_type(base_type)),
2✔
65
        patches_type,
2✔
66
        node.debug_info()
2✔
67
    );
2✔
68

69
    // Add malloc for temporary GEMM output
70
    symbolic::Expression tmp_Y_size = symbolic::mul(node.output_channels(), node.shape()[0]);
2✔
71
    for (size_t i = 0; i < dims; i++) {
5✔
72
        tmp_Y_size = symbolic::mul(tmp_Y_size, out_shape[i]);
3✔
73
    }
3✔
74
    auto tmp_Y_container = builder.find_new_name("_tmp_Y");
2✔
75
    types::Scalar tmp_Y_base_type(builder.subject().type(b.access_Y->data()).primitive_type());
2✔
76
    types::Pointer tmp_Y_type(tmp_Y_base_type);
2✔
77
    builder.add_container(tmp_Y_container, tmp_Y_type);
2✔
78
    auto [tmp_Y_malloc_block, tmp_Y_malloc_node] = stdlib::add_malloc_block(
2✔
79
        builder,
2✔
80
        new_sequence,
2✔
81
        tmp_Y_container,
2✔
82
        symbolic::mul(tmp_Y_size, symbolic::size_of_type(tmp_Y_base_type)),
2✔
83
        tmp_Y_type,
2✔
84
        node.debug_info()
2✔
85
    );
2✔
86

87
    // Add loop over batch size
88
    auto n_container = builder.find_new_name("_n");
2✔
89
    builder.add_container(n_container, indvar_type);
2✔
90
    auto n = symbolic::symbol(n_container);
2✔
91
    auto& loop_n = builder.add_map(
2✔
92
        new_sequence,
2✔
93
        n,
2✔
94
        symbolic::Lt(n, node.shape()[0]),
2✔
95
        symbolic::zero(),
2✔
96
        symbolic::add(n, symbolic::one()),
2✔
97
        ScheduleType_Sequential::create(),
2✔
98
        {},
2✔
99
        b.block->debug_info()
2✔
100
    );
2✔
101
    structured_control_flow::Sequence* current_seq = &loop_n.root();
2✔
102

103
    // Add loops over output dimensions
104
    symbolic::SymbolVec os;
2✔
105
    os.reserve(dims);
2✔
106
    for (size_t i = 0; i < dims; i++) {
5✔
107
        auto o_container = builder.find_new_name("_o");
3✔
108
        builder.add_container(o_container, indvar_type);
3✔
109
        auto o = symbolic::symbol(o_container);
3✔
110
        os.push_back(o);
3✔
111
        auto& loop_o = builder.add_map(
3✔
112
            *current_seq,
3✔
113
            o,
3✔
114
            symbolic::Lt(o, out_shape[i]),
3✔
115
            symbolic::zero(),
3✔
116
            symbolic::add(o, symbolic::one()),
3✔
117
            ScheduleType_Sequential::create(),
3✔
118
            {},
3✔
119
            b.block->debug_info()
3✔
120
        );
3✔
121
        current_seq = &loop_o.root();
3✔
122
    }
3✔
123

124
    // Add loop over channels
125
    auto c_container = builder.find_new_name("_c");
2✔
126
    builder.add_container(c_container, indvar_type);
2✔
127
    auto c = symbolic::symbol(c_container);
2✔
128
    auto& loop_c = builder.add_map(
2✔
129
        *current_seq,
2✔
130
        c,
2✔
131
        symbolic::Lt(c, node.shape()[1]),
2✔
132
        symbolic::zero(),
2✔
133
        symbolic::add(c, symbolic::one()),
2✔
134
        ScheduleType_Sequential::create(),
2✔
135
        {},
2✔
136
        b.block->debug_info()
2✔
137
    );
2✔
138
    current_seq = &loop_c.root();
2✔
139

140
    // Add loops over kernel shape
141
    symbolic::SymbolVec ks;
2✔
142
    ks.reserve(dims);
2✔
143
    for (size_t i = 0; i < dims; i++) {
5✔
144
        auto k_container = builder.find_new_name("_k");
3✔
145
        builder.add_container(k_container, indvar_type);
3✔
146
        auto k = symbolic::symbol(k_container);
3✔
147
        ks.push_back(k);
3✔
148
        auto& loop_k = builder.add_map(
3✔
149
            *current_seq,
3✔
150
            k,
3✔
151
            symbolic::Lt(k, node.kernel_shape()[i]),
3✔
152
            symbolic::zero(),
3✔
153
            symbolic::add(k, symbolic::one()),
3✔
154
            ScheduleType_Sequential::create(),
3✔
155
            {},
3✔
156
            b.block->debug_info()
3✔
157
        );
3✔
158
        current_seq = &loop_k.root();
3✔
159
    }
3✔
160

161
    // Add if/else to stay in bounds for copying
162
    symbolic::MultiExpression is;
2✔
163
    is.reserve(dims);
2✔
164
    symbolic::Condition copy_condition = symbolic::__true__();
2✔
165
    symbolic::Condition zero_condition = symbolic::__false__();
2✔
166
    for (size_t i = 0; i < dims; i++) {
5✔
167
        auto i_expr = symbolic::
3✔
168
            add(symbolic::sub(symbolic::mul(os[i], node.strides()[i]), node.pads()[i]),
3✔
169
                symbolic::mul(ks[i], node.dilations()[i]));
3✔
170
        is.push_back(i_expr);
3✔
171
        copy_condition = symbolic::
3✔
172
            And(copy_condition,
3✔
173
                symbolic::And(symbolic::Lt(i_expr, node.shape()[i + 2]), symbolic::Ge(i_expr, symbolic::zero())));
3✔
174
        zero_condition = symbolic::
3✔
175
            Or(zero_condition,
3✔
176
               symbolic::Or(symbolic::Ge(i_expr, node.shape()[i + 2]), symbolic::Lt(i_expr, symbolic::zero())));
3✔
177
    }
3✔
178
    auto& branch = builder.add_if_else(*current_seq, {}, b.block->debug_info());
2✔
179
    auto& copy_case = builder.add_case(branch, copy_condition, b.block->debug_info());
2✔
180
    auto& zero_case = builder.add_case(branch, zero_condition, b.block->debug_info());
2✔
181

182
    // Determine patches subset & tensor type
183
    data_flow::Subset patches_subset;
2✔
184
    patches_subset.push_back(n);
2✔
185
    patches_subset.insert(patches_subset.end(), os.begin(), os.end());
2✔
186
    patches_subset.push_back(c);
2✔
187
    patches_subset.insert(patches_subset.end(), ks.begin(), ks.end());
2✔
188
    symbolic::MultiExpression patches_shape;
2✔
189
    patches_shape.push_back(node.shape()[0]);
2✔
190
    patches_shape.insert(patches_shape.end(), out_shape.begin(), out_shape.end());
2✔
191
    patches_shape.push_back(node.shape()[1]);
2✔
192
    patches_shape.insert(patches_shape.end(), node.kernel_shape().begin(), node.kernel_shape().end());
2✔
193
    types::Tensor patches_tensor_type(base_type, patches_shape);
2✔
194

195
    // Determine subset for X
196
    data_flow::Subset subset_X;
2✔
197
    subset_X.push_back(n);
2✔
198
    subset_X.push_back(c);
2✔
199
    subset_X.insert(subset_X.end(), is.begin(), is.end());
2✔
200

201
    // Add copy from X to patches
202
    auto& copy_block = builder.add_block(copy_case, {}, b.block->debug_info());
2✔
203
    {
2✔
204
        auto& X_access = builder.add_access(copy_block, b.access_X->data(), b.access_X->debug_info());
2✔
205
        auto& patches_access = builder.add_access(copy_block, patches_container, node.debug_info());
2✔
206
        auto& tasklet =
2✔
207
            builder.add_tasklet(copy_block, data_flow::TaskletCode::assign, "_out", {"_in"}, node.debug_info());
2✔
208
        builder.add_computational_memlet(
2✔
209
            copy_block, X_access, tasklet, "_in", subset_X, b.iedge_X->base_type(), b.iedge_X->debug_info()
2✔
210
        );
2✔
211
        builder.add_computational_memlet(
2✔
212
            copy_block, tasklet, "_out", patches_access, patches_subset, patches_tensor_type, node.debug_info()
2✔
213
        );
2✔
214
    }
2✔
215

216
    // Add zero assignment to patches
217
    auto& zero_block = builder.add_block(zero_case, {}, b.block->debug_info());
2✔
218
    {
2✔
219
        auto& constant_zero = builder.add_constant(zero_block, "0.0", base_type, node.debug_info());
2✔
220
        auto& patches_access = builder.add_access(zero_block, patches_container, node.debug_info());
2✔
221
        auto& tasklet =
2✔
222
            builder.add_tasklet(zero_block, data_flow::TaskletCode::assign, "_out", {"_in"}, node.debug_info());
2✔
223
        builder.add_computational_memlet(zero_block, constant_zero, tasklet, "_in", {}, base_type, node.debug_info());
2✔
224
        builder.add_computational_memlet(
2✔
225
            zero_block, tasklet, "_out", patches_access, patches_subset, patches_tensor_type, node.debug_info()
2✔
226
        );
2✔
227
    }
2✔
228

229
    // Add GEMM node
230
    auto& gemm_block = builder.add_block(new_sequence, {}, b.block->debug_info());
2✔
231
    {
2✔
232
        auto& alpha = builder.add_constant(gemm_block, "1.0", base_type, node.debug_info());
2✔
233
        auto& beta = builder.add_constant(gemm_block, "0.0", base_type, node.debug_info());
2✔
234
        symbolic::Expression gemm_m = node.output_channels();
2✔
235
        symbolic::Expression gemm_n = node.shape()[0];
2✔
236
        symbolic::Expression gemm_k = node.shape()[1];
2✔
237
        for (size_t i = 0; i < dims; i++) {
5✔
238
            gemm_n = symbolic::mul(gemm_n, out_shape[i]);
3✔
239
            gemm_k = symbolic::mul(gemm_k, node.kernel_shape()[i]);
3✔
240
        }
3✔
241
        auto& libnode = math::blas::add_gemm_node(
2✔
242
            builder,
2✔
243
            gemm_block,
2✔
244
            b.access_W->data(),
2✔
245
            patches_container,
2✔
246
            tmp_Y_container,
2✔
247
            alpha,
2✔
248
            beta,
2✔
249
            precision,
2✔
250
            math::blas::BLAS_Layout::RowMajor, // layout
2✔
251
            math::blas::BLAS_Transpose::No, // transA
2✔
252
            math::blas::BLAS_Transpose::Trans, // transB
2✔
253
            gemm_m, // m
2✔
254
            gemm_n, // n
2✔
255
            gemm_k, // k
2✔
256
            gemm_k, // lda
2✔
257
            gemm_k, // ldb
2✔
258
            gemm_n, // ldc
2✔
259
            types::Pointer(types::Scalar(b.iedge_W->base_type().primitive_type())),
2✔
260
            patches_type,
2✔
261
            tmp_Y_type,
2✔
262
            base_type,
2✔
263
            node.debug_info(),
2✔
264
            b.access_W->debug_info(),
2✔
265
            node.debug_info(),
2✔
266
            b.access_Y->debug_info(),
2✔
267
            b.iedge_W->debug_info(),
2✔
268
            node.debug_info(),
2✔
269
            b.iedge_Y->debug_info(),
2✔
270
            math::blas::ImplementationType_BLAS
2✔
271
        );
2✔
272
    }
2✔
273

274
    // Add loop over batch size (again)
275
    auto& loop_n_2 = builder.add_map(
2✔
276
        new_sequence,
2✔
277
        n,
2✔
278
        symbolic::Lt(n, node.shape()[0]),
2✔
279
        symbolic::zero(),
2✔
280
        symbolic::add(n, symbolic::one()),
2✔
281
        ScheduleType_Sequential::create(),
2✔
282
        {},
2✔
283
        b.block->debug_info()
2✔
284
    );
2✔
285
    current_seq = &loop_n_2.root();
2✔
286

287
    // Add loop over output channels
288
    auto l_container = builder.find_new_name("_l");
2✔
289
    builder.add_container(l_container, indvar_type);
2✔
290
    auto l = symbolic::symbol(l_container);
2✔
291
    auto& loop_l = builder.add_map(
2✔
292
        *current_seq,
2✔
293
        l,
2✔
294
        symbolic::Lt(l, node.output_channels()),
2✔
295
        symbolic::zero(),
2✔
296
        symbolic::add(l, symbolic::one()),
2✔
297
        ScheduleType_Sequential::create(),
2✔
298
        {},
2✔
299
        b.block->debug_info()
2✔
300
    );
2✔
301
    current_seq = &loop_l.root();
2✔
302

303
    // Add loops over output dimensions (again)
304
    for (size_t i = 0; i < dims; i++) {
5✔
305
        auto o_container = builder.find_new_name("_o");
3✔
306
        builder.add_container(o_container, indvar_type);
3✔
307
        auto o = symbolic::symbol(o_container);
3✔
308
        auto& loop_o = builder.add_map(
3✔
309
            *current_seq,
3✔
310
            o,
3✔
311
            symbolic::Lt(o, out_shape[i]),
3✔
312
            symbolic::zero(),
3✔
313
            symbolic::add(o, symbolic::one()),
3✔
314
            ScheduleType_Sequential::create(),
3✔
315
            {},
3✔
316
            b.block->debug_info()
3✔
317
        );
3✔
318
        current_seq = &loop_o.root();
3✔
319
        os[i] = o;
3✔
320
    }
3✔
321

322
    // Add transposed copy from temporary GEMM output to Y + add bias if available
323
    data_flow::Subset tmp_Y_subset;
2✔
324
    tmp_Y_subset.push_back(l);
2✔
325
    tmp_Y_subset.push_back(n);
2✔
326
    tmp_Y_subset.insert(tmp_Y_subset.end(), os.begin(), os.end());
2✔
327
    symbolic::MultiExpression tmp_Y_shape;
2✔
328
    tmp_Y_shape.push_back(node.output_channels());
2✔
329
    tmp_Y_shape.push_back(node.shape()[0]);
2✔
330
    tmp_Y_shape.insert(tmp_Y_shape.end(), out_shape.begin(), out_shape.end());
2✔
331
    types::Tensor tmp_Y_tensor_type(tmp_Y_base_type, tmp_Y_shape);
2✔
332
    data_flow::Subset Y_subset;
2✔
333
    Y_subset.push_back(n);
2✔
334
    Y_subset.push_back(l);
2✔
335
    Y_subset.insert(Y_subset.end(), os.begin(), os.end());
2✔
336
    auto& transpose_block = builder.add_block(*current_seq, {}, b.block->debug_info());
2✔
337
    if (b.has_bias) {
2✔
NEW
338
        auto& tmp_Y_access = builder.add_access(transpose_block, tmp_Y_container, node.debug_info());
×
NEW
339
        auto& B_access = builder.add_access(transpose_block, b.access_B->data(), b.access_B->debug_info());
×
NEW
340
        auto& Y_access = builder.add_access(transpose_block, b.access_Y->data(), b.access_Y->debug_info());
×
NEW
341
        auto& tasklet =
×
342
            builder
×
NEW
343
                .add_tasklet(transpose_block, data_flow::TaskletCode::fp_add, "_out", {"_in1", "_in2"}, node.debug_info());
×
NEW
344
        builder.add_computational_memlet(
×
NEW
345
            transpose_block, tmp_Y_access, tasklet, "_in1", tmp_Y_subset, tmp_Y_tensor_type, node.debug_info()
×
NEW
346
        );
×
NEW
347
        builder.add_computational_memlet(
×
NEW
348
            transpose_block, B_access, tasklet, "_in2", {l}, b.iedge_B->base_type(), b.iedge_B->debug_info()
×
NEW
349
        );
×
NEW
350
        builder.add_computational_memlet(
×
NEW
351
            transpose_block, tasklet, "_out", Y_access, Y_subset, b.iedge_Y->base_type(), b.iedge_Y->debug_info()
×
NEW
352
        );
×
353
    } else {
2✔
354
        auto& tmp_Y_access = builder.add_access(transpose_block, tmp_Y_container, node.debug_info());
2✔
355
        auto& Y_access = builder.add_access(transpose_block, b.access_Y->data(), b.access_Y->debug_info());
2✔
356
        auto& tasklet =
2✔
357
            builder.add_tasklet(transpose_block, data_flow::TaskletCode::assign, "_out", {"_in"}, node.debug_info());
2✔
358
        builder.add_computational_memlet(
2✔
359
            transpose_block, tmp_Y_access, tasklet, "_in", tmp_Y_subset, tmp_Y_tensor_type, node.debug_info()
2✔
360
        );
2✔
361
        builder.add_computational_memlet(
2✔
362
            transpose_block, tasklet, "_out", Y_access, Y_subset, b.iedge_Y->base_type(), b.iedge_Y->debug_info()
2✔
363
        );
2✔
364
    }
2✔
365

366
    // Add free for patches container
367
    auto [patches_free_block, patches_free_node] =
2✔
368
        stdlib::add_free_block(builder, new_sequence, patches_container, patches_type, node.debug_info());
2✔
369

370
    // Add free for temporary GEMM output
371
    auto [tmp_Y_free_block, tmp_Y_free_node] =
2✔
372
        stdlib::add_free_block(builder, new_sequence, tmp_Y_container, tmp_Y_type, node.debug_info());
2✔
373

374
    // Clean up the original block
375
    builder.clear_code_node_legacy(*b.block, node);
2✔
376
    builder.remove_child(*b.block_parent, b.block_index + 1);
2✔
377

378
    return true;
2✔
379
}
3✔
380
} // namespace offloading
381
} // namespace sdfg
STATUS · Troubleshooting · Open an Issue · Sales · Support · CAREERS · ENTERPRISE · START FREE TRIAL · SCHEDULE DEMO
ANNOUNCEMENTS · TWITTER · TOS & SLA · Supported CI Services · What's a CI service? · Automated Testing

© 2026 Coveralls, Inc