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

daisytuner / docc / 30543265435

30 Jul 2026 12:36PM UTC coverage: 64.549% (-0.2%) from 64.738%
30543265435

Pull #913

github

web-flow
Merge 6970d8cbe into 4a519a416
Pull Request #913: delegates instrumentation event type decision to specific node dispat…

17 of 236 new or added lines in 25 files covered. (7.2%)

7 existing lines in 7 files now uncovered.

44347 of 68703 relevant lines covered (64.55%)

714.95 hits per line

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

41.65
/sdfg/src/data_flow/library_nodes/math/blas/gemm_node.cpp
1
#include "sdfg/data_flow/library_nodes/math/blas/gemm_node.h"
2

3
#include "sdfg/analysis/analysis.h"
4
#include "sdfg/builder/structured_sdfg_builder.h"
5

6
namespace sdfg {
7
namespace math {
8
namespace blas {
9

10
GEMMNode::GEMMNode(
11
    size_t element_id,
12
    const DebugInfo& debug_info,
13
    const graph::Vertex vertex,
14
    data_flow::DataFlowGraph& parent,
15
    const data_flow::ImplementationType& implementation_type,
16
    const BLAS_Precision& precision,
17
    const BLAS_Layout& layout,
18
    const BLAS_Transpose& trans_a,
19
    const BLAS_Transpose& trans_b,
20
    symbolic::Expression m,
21
    symbolic::Expression n,
22
    symbolic::Expression k,
23
    symbolic::Expression lda,
24
    symbolic::Expression ldb,
25
    symbolic::Expression ldc
26
)
27
    : BLASNode(
35✔
28
          element_id,
35✔
29
          debug_info,
35✔
30
          vertex,
35✔
31
          parent,
35✔
32
          LibraryNodeType_GEMM,
35✔
33
          {},
35✔
34
          {"__A", "__B", "__C", "__alpha", "__beta"},
35✔
35
          implementation_type,
35✔
36
          precision
35✔
37
      ),
35✔
38
      layout_(layout), trans_a_(trans_a), trans_b_(trans_b), m_(m), n_(n), k_(k), lda_(lda), ldb_(ldb), ldc_(ldc) {}
35✔
39

40
BLAS_Layout GEMMNode::layout() const { return this->layout_; };
4✔
41

42
BLAS_Transpose GEMMNode::trans_a() const { return this->trans_a_; };
5✔
43

44
BLAS_Transpose GEMMNode::trans_b() const { return this->trans_b_; };
5✔
45

46
symbolic::Expression GEMMNode::m() const { return this->m_; };
19✔
47

48
symbolic::Expression GEMMNode::n() const { return this->n_; };
21✔
49

50
symbolic::Expression GEMMNode::k() const { return this->k_; };
19✔
51

52
symbolic::Expression GEMMNode::lda() const { return this->lda_; };
11✔
53

54
symbolic::Expression GEMMNode::ldb() const { return this->ldb_; };
11✔
55

56
symbolic::Expression GEMMNode::ldc() const { return this->ldc_; };
11✔
57

58
symbolic::SymbolSet GEMMNode::symbols() const {
×
59
    symbolic::SymbolSet syms;
×
60

61
    for (auto& atom : symbolic::atoms(this->m_)) {
×
62
        syms.insert(atom);
×
63
    }
×
64
    for (auto& atom : symbolic::atoms(this->n_)) {
×
65
        syms.insert(atom);
×
66
    }
×
67
    for (auto& atom : symbolic::atoms(this->k_)) {
×
68
        syms.insert(atom);
×
69
    }
×
70
    for (auto& atom : symbolic::atoms(this->lda_)) {
×
71
        syms.insert(atom);
×
72
    }
×
73
    for (auto& atom : symbolic::atoms(this->ldb_)) {
×
74
        syms.insert(atom);
×
75
    }
×
76
    for (auto& atom : symbolic::atoms(this->ldc_)) {
×
77
        syms.insert(atom);
×
78
    }
×
79

80
    return syms;
×
81
};
×
82

83
void GEMMNode::replace(const symbolic::Expression old_expression, const symbolic::Expression new_expression) {
×
84
    this->m_ = symbolic::subs(this->m_, old_expression, new_expression);
×
85
    this->n_ = symbolic::subs(this->n_, old_expression, new_expression);
×
86
    this->k_ = symbolic::subs(this->k_, old_expression, new_expression);
×
87
    this->lda_ = symbolic::subs(this->lda_, old_expression, new_expression);
×
88
    this->ldb_ = symbolic::subs(this->ldb_, old_expression, new_expression);
×
89
    this->ldc_ = symbolic::subs(this->ldc_, old_expression, new_expression);
×
90
};
×
91

92
void GEMMNode::replace(const symbolic::ExpressionMapping& replacements) {
×
93
    this->m_ = symbolic::subs(this->m_, replacements);
×
94
    this->n_ = symbolic::subs(this->n_, replacements);
×
95
    this->k_ = symbolic::subs(this->k_, replacements);
×
96
    this->lda_ = symbolic::subs(this->lda_, replacements);
×
97
    this->ldb_ = symbolic::subs(this->ldb_, replacements);
×
98
    this->ldc_ = symbolic::subs(this->ldc_, replacements);
×
99
};
×
100

101
void GEMMNode::validate(const Function& function) const { BLASNode::validate(function); }
14✔
102

103
passes::LibNodeExpander::ExpandOutcome GEMMNode::
104
    expand(passes::LibNodeExpander::ExpandContext& context, structured_control_flow::Block& block) {
8✔
105
    auto& dataflow = this->get_parent();
8✔
106

107
    if (trans_a_ == BLAS_Transpose::ConjTrans || trans_b_ == BLAS_Transpose::ConjTrans) {
8✔
108
        return context.unable();
×
109
    }
×
110

111
    auto primitive_type = scalar_primitive();
8✔
112
    if (primitive_type == types::PrimitiveType::Void) {
8✔
113
        return context.unable();
×
114
    }
×
115

116
    types::Scalar scalar_type(primitive_type);
8✔
117

118
    auto in_edges = dataflow.in_edges(*this);
8✔
119
    auto in_edges_it = in_edges.begin();
8✔
120

121
    data_flow::Memlet* iedge_a = nullptr;
8✔
122
    data_flow::Memlet* iedge_b = nullptr;
8✔
123
    data_flow::Memlet* iedge_c = nullptr;
8✔
124
    data_flow::Memlet* alpha_edge = nullptr;
8✔
125
    data_flow::Memlet* beta_edge = nullptr;
8✔
126
    while (in_edges_it != in_edges.end()) {
48✔
127
        auto& edge = *in_edges_it;
40✔
128
        auto dst_conn = edge.dst_conn();
40✔
129
        if (dst_conn == "__A") {
40✔
130
            iedge_a = &edge;
8✔
131
        } else if (dst_conn == "__B") {
32✔
132
            iedge_b = &edge;
8✔
133
        } else if (dst_conn == "__C") {
24✔
134
            iedge_c = &edge;
8✔
135
        } else if (dst_conn == "__alpha") {
16✔
136
            alpha_edge = &edge;
8✔
137
        } else if (dst_conn == "__beta") {
8✔
138
            beta_edge = &edge;
8✔
139
        } else {
8✔
140
            throw InvalidSDFGException("GEMMNode has unexpected input: " + dst_conn);
×
141
        }
×
142
        ++in_edges_it;
40✔
143
    }
40✔
144

145
    using Dir = passes::LibNodeExpander::InputUse;
8✔
146
    auto standalone = context.replacement_requires_access_nodes(
8✔
147
        {Dir::IndirectRead, Dir::IndirectRead, Dir::IndirectReadWrite, Dir::Scalar, Dir::Scalar}
8✔
148
    );
8✔
149

150
    if (!standalone) {
8✔
151
        return context.unable();
×
152
    }
×
153

154
    // Add new graph after the current block
155
    auto& new_sequence = standalone->replace_with_sequence();
8✔
156
    auto& builder = standalone->builder();
8✔
157

158
    // Add maps
159
    std::vector<symbolic::Expression> indvar_ends{this->m(), this->n(), this->k()};
8✔
160
    data_flow::Subset new_subset;
8✔
161
    structured_control_flow::Sequence* last_scope = &new_sequence;
8✔
162
    structured_control_flow::StructuredLoop* last_map = nullptr;
8✔
163
    structured_control_flow::StructuredLoop* output_loop = nullptr;
8✔
164
    std::vector<std::string> indvar_names{"_i", "_j", "_k"};
8✔
165

166
    std::string sum_var = builder.find_new_name("_sum");
8✔
167
    builder.add_container(sum_var, scalar_type);
8✔
168

169
    for (size_t i = 0; i < 3; i++) {
32✔
170
        auto dim_begin = symbolic::zero();
24✔
171
        auto& dim_end = indvar_ends[i];
24✔
172

173
        std::string indvar_str = builder.find_new_name(indvar_names[i]);
24✔
174
        builder.add_container(indvar_str, types::Scalar(types::PrimitiveType::UInt64));
24✔
175

176
        auto indvar = symbolic::symbol(indvar_str);
24✔
177
        auto init = dim_begin;
24✔
178
        auto update = symbolic::add(indvar, symbolic::one());
24✔
179
        auto condition = symbolic::Lt(indvar, dim_end);
24✔
180
        if (i < 2) {
24✔
181
            last_map = &builder.add_map(
16✔
182
                *last_scope,
16✔
183
                indvar,
16✔
184
                condition,
16✔
185
                init,
16✔
186
                update,
16✔
187
                structured_control_flow::ScheduleType_Sequential::create(),
16✔
188
                block.debug_info()
16✔
189
            );
16✔
190
        } else {
16✔
191
            last_map = &builder.add_for(*last_scope, indvar, condition, init, update, block.debug_info());
8✔
192
        }
8✔
193
        last_scope = &last_map->root();
24✔
194

195
        if (i == 1) {
24✔
196
            output_loop = last_map;
8✔
197
        }
8✔
198

199
        new_subset.push_back(indvar);
24✔
200
    }
24✔
201

202

203
    // Add code
204
    auto& init_block = builder.add_block_before(output_loop->root(), *last_map, block.debug_info());
8✔
205
    auto& sum_init = builder.add_access(init_block, sum_var, block.debug_info());
8✔
206

207
    auto& zero_node = builder.add_constant(init_block, "0.0", alpha_edge->base_type(), block.debug_info());
8✔
208
    auto& init_tasklet = builder.add_tasklet(init_block, data_flow::assign, "_out", {"_in"}, block.debug_info());
8✔
209
    builder.add_computational_memlet(init_block, zero_node, init_tasklet, "_in", {}, block.debug_info());
8✔
210
    builder.add_computational_memlet(init_block, init_tasklet, "_out", sum_init, {}, block.debug_info());
8✔
211

212
    auto& code_block = builder.add_block(*last_scope, block.debug_info());
8✔
213
    auto& input_node_a_new = standalone->add_indirect_read_access(code_block, A_INPUT_IDX);
8✔
214
    auto& input_node_b_new = standalone->add_indirect_read_access(code_block, B_INPUT_IDX);
8✔
215

216
    auto& core_fma =
8✔
217
        builder.add_tasklet(code_block, data_flow::fp_fma, "_out", {"_in1", "_in2", "_in3"}, block.debug_info());
8✔
218
    auto& sum_in = builder.add_access(code_block, sum_var, block.debug_info());
8✔
219
    auto& sum_out = builder.add_access(code_block, sum_var, block.debug_info());
8✔
220
    builder.add_computational_memlet(code_block, sum_in, core_fma, "_in3", {}, block.debug_info());
8✔
221

222
    // Row-major indexing: address = ld * row + col
223
    // No transpose: A is m×k, access A[i, k] => lda*i + k
224
    // Transpose:    A is k×m stored, access A[k, i] => lda*k + i
225
    symbolic::Expression a_idx = (trans_a_ == BLAS_Transpose::Trans)
8✔
226
                                     ? symbolic::add(symbolic::mul(lda(), new_subset[2]), new_subset[0])
8✔
227
                                     : symbolic::add(symbolic::mul(lda(), new_subset[0]), new_subset[2]);
8✔
228
    builder.add_computational_memlet(
8✔
229
        code_block, input_node_a_new, core_fma, "_in1", {a_idx}, iedge_a->base_type(), iedge_a->debug_info()
8✔
230
    );
8✔
231
    // No transpose: B is k×n, access B[k, j] => ldb*k + j
232
    // Transpose:    B is n×k stored, access B[j, k] => ldb*j + k
233
    symbolic::Expression b_idx = (trans_b_ == BLAS_Transpose::Trans)
8✔
234
                                     ? symbolic::add(symbolic::mul(ldb(), new_subset[1]), new_subset[2])
8✔
235
                                     : symbolic::add(symbolic::mul(ldb(), new_subset[2]), new_subset[1]);
8✔
236
    builder.add_computational_memlet(
8✔
237
        code_block, input_node_b_new, core_fma, "_in2", {b_idx}, iedge_b->base_type(), iedge_b->debug_info()
8✔
238
    );
8✔
239
    builder.add_computational_memlet(code_block, core_fma, "_out", sum_out, {}, iedge_c->debug_info());
8✔
240

241
    auto& flush_block = builder.add_block_after(output_loop->root(), *last_map, block.debug_info());
8✔
242
    auto& sum_final = builder.add_access(flush_block, sum_var, block.debug_info());
8✔
243
    auto& input_node_c_new = standalone->add_indirect_read_access(flush_block, C_INPUT_IDX);
8✔
244
    symbolic::Expression c_idx = symbolic::add(symbolic::mul(ldc(), new_subset[0]), new_subset[1]);
8✔
245

246
    auto& scale_sum_tasklet =
8✔
247
        builder.add_tasklet(flush_block, data_flow::TaskletCode::fp_mul, "_out", {"_in1", "_in2"}, block.debug_info());
8✔
248
    builder.add_computational_memlet(flush_block, sum_final, scale_sum_tasklet, "_in1", {}, block.debug_info());
8✔
249
    auto& alpha_node = standalone->add_scalar_input_access(flush_block, ALPHA_INPUT_IDX);
8✔
250
    builder.add_computational_memlet(flush_block, alpha_node, scale_sum_tasklet, "_in2", {}, block.debug_info());
8✔
251

252
    std::string scaled_sum_temp = builder.find_new_name("scaled_sum_temp");
8✔
253
    builder.add_container(scaled_sum_temp, scalar_type);
8✔
254
    auto& scaled_sum_final = builder.add_access(flush_block, scaled_sum_temp, block.debug_info());
8✔
255
    builder.add_computational_memlet(
8✔
256
        flush_block, scale_sum_tasklet, "_out", scaled_sum_final, {}, scalar_type, block.debug_info()
8✔
257
    );
8✔
258

259
    auto& scale_input_tasklet =
8✔
260
        builder.add_tasklet(flush_block, data_flow::TaskletCode::fp_mul, "_out", {"_in1", "_in2"}, block.debug_info());
8✔
261
    builder.add_computational_memlet(
8✔
262
        flush_block, input_node_c_new, scale_input_tasklet, "_in1", {c_idx}, iedge_c->base_type(), iedge_c->debug_info()
8✔
263
    );
8✔
264
    auto& beta_node = standalone->add_scalar_input_access(flush_block, BETA_INPUT_IDX);
8✔
265
    builder.add_computational_memlet(flush_block, beta_node, scale_input_tasklet, "_in2", {}, block.debug_info());
8✔
266

267
    std::string scaled_input_temp = builder.find_new_name("scaled_input_temp");
8✔
268
    builder.add_container(scaled_input_temp, scalar_type);
8✔
269
    auto& scaled_input_c = builder.add_access(flush_block, scaled_input_temp, block.debug_info());
8✔
270
    builder.add_computational_memlet(
8✔
271
        flush_block, scale_input_tasklet, "_out", scaled_input_c, {}, scalar_type, block.debug_info()
8✔
272
    );
8✔
273

274
    auto& flush_add_tasklet =
8✔
275
        builder.add_tasklet(flush_block, data_flow::TaskletCode::fp_add, "_out", {"_in1", "_in2"}, block.debug_info());
8✔
276
    auto& output_node_new = standalone->add_indirect_write_access(flush_block, C_INPUT_IDX);
8✔
277
    builder.add_computational_memlet(
8✔
278
        flush_block, scaled_sum_final, flush_add_tasklet, "_in1", {}, scalar_type, block.debug_info()
8✔
279
    );
8✔
280
    builder.add_computational_memlet(
8✔
281
        flush_block, scaled_input_c, flush_add_tasklet, "_in2", {}, scalar_type, block.debug_info()
8✔
282
    );
8✔
283
    builder.add_computational_memlet(
8✔
284
        flush_block, flush_add_tasklet, "_out", output_node_new, {c_idx}, iedge_c->base_type(), iedge_c->debug_info()
8✔
285
    );
8✔
286

287
    return standalone->successfully_expanded();
8✔
288
}
8✔
289

290
symbolic::Expression GEMMNode::flop() const {
×
291
    return flops(symbolic::__true__(), symbolic::__true__(), symbolic::__true__(), symbolic::__true__());
×
292
}
×
293

294
symbolic::Expression GEMMNode::flops(
295
    symbolic::Condition alpha_non_zero,
296
    symbolic::Condition alpha_non_ident,
297
    symbolic::Condition beta_non_zero,
298
    symbolic::Condition beta_non_ident
299
) const {
×
300
    auto res_elems = symbolic::mul(this->m_, this->n_);
×
301

302
    // conditional on alpha != 0.0
303
    auto mm_mul_ops = symbolic::mul(symbolic::mul(res_elems, this->k_), alpha_non_zero);
×
304
    auto mm_sum_ops = symbolic::mul(symbolic::mul(res_elems, symbolic::sub(this->k_, symbolic::one())), alpha_non_zero);
×
305
    // conditional on alpha != 1.0 && alpha != 0.0
306
    auto mm_alpha_scale_ops = symbolic::mul(res_elems, symbolic::And(alpha_non_ident, alpha_non_zero));
×
307
    // conditional on beta != 1.0 && beta != 0.0
308
    auto mm_beta_scale_ops = symbolic::mul(res_elems, symbolic::And(beta_non_ident, beta_non_zero));
×
309
    auto mm_beta_scaled_sum_ops = symbolic::mul(res_elems, beta_non_zero);
×
310
    auto mul_ops = symbolic::add(mm_mul_ops, symbolic::add(mm_alpha_scale_ops, mm_beta_scale_ops));
×
311
    auto add_ops = symbolic::add(mm_sum_ops, mm_beta_scaled_sum_ops);
×
312
    return symbolic::add(mul_ops, add_ops);
×
313
}
×
314

315
std::unique_ptr<data_flow::DataFlowNode> GEMMNode::
316
    clone(size_t element_id, const graph::Vertex vertex, data_flow::DataFlowGraph& parent) const {
×
317
    auto node_clone = std::unique_ptr<GEMMNode>(new GEMMNode(
×
318
        element_id,
×
319
        this->debug_info(),
×
320
        vertex,
×
321
        parent,
×
322
        this->implementation_type_,
×
323
        this->precision_,
×
324
        this->layout_,
×
325
        this->trans_a_,
×
326
        this->trans_b_,
×
327
        this->m_,
×
328
        this->n_,
×
329
        this->k_,
×
330
        this->lda_,
×
331
        this->ldb_,
×
332
        this->ldc_
×
333
    ));
×
334
    return std::move(node_clone);
×
335
}
×
336

337
std::string GEMMNode::toStr() const {
×
338
    return LibraryNode::toStr() + "(" + static_cast<char>(precision_) + ", " +
×
339
           std::string(BLAS_Layout_to_short_string(layout_)) + ", " + BLAS_Transpose_to_char(trans_a_) +
×
340
           BLAS_Transpose_to_char(trans_b_) + ", " + m_->__str__() + ", " + n_->__str__() + ", " + k_->__str__() +
×
341
           ", " + lda_->__str__() + ", " + ldb_->__str__() + ", " + ldc_->__str__() + ")";
×
342
}
×
343

344
symbolic::Expression GEMMNode::calc_matrix_access_range(
345
    const symbolic::Expression& outer_dim,
346
    const symbolic::Expression& inner_dim,
347
    const symbolic::Expression& line_size,
348
    BLAS_Transpose trans,
349
    BLAS_Layout layout
350
) {
×
351
    if ((trans == BLAS_Transpose::No) ^ (layout == BLAS_Layout::ColMajor)) {
×
352
        return symbolic::mul(outer_dim, line_size);
×
353
    } else {
×
354
        return symbolic::mul(inner_dim, line_size);
×
355
    }
×
356
}
×
357

358

359
data_flow::PointerAccessType GEMMNode::pointer_access_type(int input_idx) const {
×
360
    if (input_idx == 0) { // A: m x k
×
361
        return data_flow::PointerAccessMeta::
×
362
            create_read_only(calc_matrix_access_range(m_, k_, lda_, trans_a_, layout_), true);
×
363
    } else if (input_idx == 1) { // B: k x n
×
364
        return data_flow::PointerAccessMeta::
×
365
            create_read_only(calc_matrix_access_range(k_, n_, ldb_, trans_b_, layout_), true);
×
366
    } else if (input_idx == 2) {
×
367
        // for beta == 0, there would no reads of C. But we currently have no mechanism to access const-prop knowledge
368
        // like tha
369
        if (symbolic::eq(ldc_, n_)) { // non-sparse access over the m x n range
×
370
            return data_flow::PointerAccessMeta::
×
371
                create_full_write_only(calc_matrix_access_range(m_, n_, ldc_, BLAS_Transpose::No, layout_), true);
×
372
        } else {
×
373
            // sparse access. But with only Convex Pattern for now, we cannot represent which values are
374
            auto pattern =
×
375
                data_flow::ConvexAccessPattern::create(calc_matrix_access_range(m_, n_, ldc_, BLAS_Transpose::No, layout_)
×
376
                );
×
377
            // full-overwritten and which are DC.
378
            return data_flow::PointerAccessMeta::create_generic(pattern->ref(), std::move(pattern), true);
×
379
        }
×
380
    } else {
×
381
        return LibraryNode::pointer_access_type(input_idx);
×
382
    }
×
383
}
×
384

385
nlohmann::json GEMMNodeSerializer::serialize(const data_flow::LibraryNode& library_node) {
×
386
    const GEMMNode& gemm_node = static_cast<const GEMMNode&>(library_node);
×
387
    nlohmann::json j;
×
388

389
    serializer::JSONSerializer serializer;
×
390
    j["code"] = gemm_node.code().value();
×
391
    j["precision"] = gemm_node.precision();
×
392
    j["layout"] = gemm_node.layout();
×
393
    j["trans_a"] = gemm_node.trans_a();
×
394
    j["trans_b"] = gemm_node.trans_b();
×
395
    j["m"] = serializer.expression(gemm_node.m());
×
396
    j["n"] = serializer.expression(gemm_node.n());
×
397
    j["k"] = serializer.expression(gemm_node.k());
×
398
    j["lda"] = serializer.expression(gemm_node.lda());
×
399
    j["ldb"] = serializer.expression(gemm_node.ldb());
×
400
    j["ldc"] = serializer.expression(gemm_node.ldc());
×
401

402
    return j;
×
403
}
×
404

405
data_flow::LibraryNode& GEMMNodeSerializer::deserialize(
406
    const nlohmann::json& j, builder::StructuredSDFGBuilder& builder, structured_control_flow::Block& parent
407
) {
×
408
    // Assertions for required fields
409
    assert(j.contains("element_id"));
×
410
    assert(j.contains("code"));
×
411
    assert(j.contains("debug_info"));
×
412

413
    auto code = j["code"].get<std::string>();
×
414
    if (code != LibraryNodeType_GEMM.value()) {
×
415
        throw std::runtime_error("Invalid library node code");
×
416
    }
×
417

418
    // Extract debug info using JSONSerializer
419
    sdfg::serializer::JSONSerializer serializer;
×
420
    DebugInfo debug_info = serializer.json_to_debug_info(j["debug_info"]);
×
421

422
    auto precision = j.at("precision").get<BLAS_Precision>();
×
423
    auto layout = j.at("layout").get<BLAS_Layout>();
×
424
    auto trans_a = j.at("trans_a").get<BLAS_Transpose>();
×
425
    auto trans_b = j.at("trans_b").get<BLAS_Transpose>();
×
426
    auto m = symbolic::parse(j.at("m"));
×
427
    auto n = symbolic::parse(j.at("n"));
×
428
    auto k = symbolic::parse(j.at("k"));
×
429
    auto lda = symbolic::parse(j.at("lda"));
×
430
    auto ldb = symbolic::parse(j.at("ldb"));
×
431
    auto ldc = symbolic::parse(j.at("ldc"));
×
432

433
    auto implementation_type = j.at("implementation_type").get<std::string>();
×
434

435
    return builder.add_library_node<
×
436
        GEMMNode>(parent, debug_info, implementation_type, precision, layout, trans_a, trans_b, m, n, k, lda, ldb, ldc);
×
437
}
×
438

439
GEMMNodeDispatcher_BLAS::GEMMNodeDispatcher_BLAS(
440
    codegen::LanguageExtension& language_extension,
441
    const Function& function,
442
    const data_flow::DataFlowGraph& data_flow_graph,
443
    const GEMMNode& node
444
)
445
    : codegen::LibraryNodeDispatcher(language_extension, function, data_flow_graph, node) {}
×
446

447
void GEMMNodeDispatcher_BLAS::dispatch_code_with_edges(
448
    codegen::CodegenOutput& out,
449
    std::vector<codegen::DispatchInput>& inputs,
450
    std::vector<codegen::DispatchOutput>& outputs
451
) {
×
452
    auto& gemm_node = static_cast<const GEMMNode&>(this->node_);
×
453

454
    sdfg::types::Scalar base_type(types::PrimitiveType::Void);
×
455
    switch (gemm_node.precision()) {
×
456
        case BLAS_Precision::h:
×
457
            base_type = types::Scalar(types::PrimitiveType::Half);
×
458
            break;
×
459
        case BLAS_Precision::s:
×
460
            base_type = types::Scalar(types::PrimitiveType::Float);
×
461
            break;
×
462
        case BLAS_Precision::d:
×
463
            base_type = types::Scalar(types::PrimitiveType::Double);
×
464
            break;
×
465
        default:
×
466
            throw std::runtime_error("Invalid BLAS_Precision value");
×
467
    }
×
468

469
    out.library_snippet_factory.require_dependency(BLASLibDependency::instance());
×
470

471
    out.stream << "cblas_" << BLAS_Precision_to_string(gemm_node.precision()) << "gemm(";
×
472
    out.stream.changeIndent(+4);
×
473
    out.stream << BLAS_Layout_to_string(gemm_node.layout());
×
474
    out.stream << ", ";
×
475
    out.stream << BLAS_Transpose_to_string(gemm_node.trans_a());
×
476
    out.stream << ", ";
×
477
    out.stream << BLAS_Transpose_to_string(gemm_node.trans_b());
×
478
    out.stream << ", ";
×
479
    out.stream << this->language_extension_.expression(gemm_node.m());
×
480
    out.stream << ", ";
×
481
    out.stream << this->language_extension_.expression(gemm_node.n());
×
482
    out.stream << ", ";
×
483
    out.stream << this->language_extension_.expression(gemm_node.k());
×
484
    out.stream << ", ";
×
485
    out.stream << inputs.at(GEMMNode::ALPHA_INPUT_IDX).expr;
×
486
    out.stream << ", ";
×
487
    out.stream << inputs.at(GEMMNode::A_INPUT_IDX).expr;
×
488
    out.stream << ", ";
×
489
    out.stream << this->language_extension_.expression(gemm_node.lda());
×
490
    out.stream << ", ";
×
491
    out.stream << inputs.at(GEMMNode::B_INPUT_IDX).expr;
×
492
    out.stream << ", ";
×
493
    out.stream << this->language_extension_.expression(gemm_node.ldb());
×
494
    out.stream << ", ";
×
495
    out.stream << inputs.at(GEMMNode::BETA_INPUT_IDX).expr;
×
496
    out.stream << ", ";
×
497
    out.stream << inputs.at(GEMMNode::C_INPUT_IDX).expr;
×
498
    out.stream << ", ";
×
499
    out.stream << this->language_extension_.expression(gemm_node.ldc());
×
500

501
    out.stream.changeIndent(-4);
×
502
    out.stream << ");" << std::endl;
×
503
}
×
504

NEW
505
codegen::InstrumentationInfo GEMMNodeDispatcher_BLAS::instrumentation_info() const {
×
NEW
506
    return {
×
NEW
507
        node_.element_id(),
×
NEW
508
        std::string(node_.element_type()) + ":::" + node_.code().value(),
×
NEW
509
        codegen::TargetType_CPU_PARALLEL,
×
NEW
510
        codegen::InstrumentationEventType::CPU,
×
NEW
511
        analysis::LoopInfo{},
×
NEW
512
        {}
×
NEW
513
    };
×
NEW
514
}
×
515

516
GEMMNode& add_gemm_node(
517
    builder::StructuredSDFGBuilder& builder,
518
    Block& block,
519
    const std::string& ptr_a,
520
    const std::string& ptr_b,
521
    const std::string& ptr_c,
522
    data_flow::AccessNode& alpha_node,
523
    data_flow::AccessNode& beta_node,
524
    const BLAS_Precision& precision,
525
    const BLAS_Layout& layout,
526
    const BLAS_Transpose& trans_a,
527
    const BLAS_Transpose& trans_b,
528
    symbolic::Expression& m,
529
    symbolic::Expression& n,
530
    symbolic::Expression& k,
531
    symbolic::Expression& lda,
532
    symbolic::Expression& ldb,
533
    symbolic::Expression& ldc,
534
    const types::IType& a_type,
535
    const types::IType& b_type,
536
    const types::IType& c_type,
537
    const types::IType& factor_type,
538
    DebugInfo debug_info,
539
    DebugInfo a_access_deb_info,
540
    DebugInfo b_access_deb_info,
541
    DebugInfo c_access_deb_info,
542
    DebugInfo a_edge_deb_info,
543
    DebugInfo b_edge_deb_info,
544
    DebugInfo c_edge_deb_info,
545
    data_flow::ImplementationType impl_type
546
) {
6✔
547
    auto& gemm_node = builder.add_library_node<sdfg::math::blas::GEMMNode>(
6✔
548
        block, debug_info, std::move(impl_type), precision, layout, trans_a, trans_b, m, n, k, lda, ldb, ldc
6✔
549
    );
6✔
550

551
    // Add access nodes
552
    auto& a_node_in = builder.add_access(block, ptr_a, a_access_deb_info);
6✔
553
    auto& b_node_in = builder.add_access(block, ptr_b, b_access_deb_info);
6✔
554
    auto& c_node_in = builder.add_access(block, ptr_c, c_access_deb_info);
6✔
555

556
    // Add edges
557
    builder.add_computational_memlet(block, a_node_in, gemm_node, "__A", {}, a_type, a_edge_deb_info);
6✔
558
    builder.add_computational_memlet(block, b_node_in, gemm_node, "__B", {}, b_type, b_edge_deb_info);
6✔
559
    builder.add_computational_memlet(block, c_node_in, gemm_node, "__C", {}, c_type, c_edge_deb_info);
6✔
560
    builder.add_computational_memlet(block, alpha_node, gemm_node, "__alpha", {}, factor_type, debug_info);
6✔
561
    builder.add_computational_memlet(block, beta_node, gemm_node, "__beta", {}, factor_type, debug_info);
6✔
562

563
    return static_cast<GEMMNode&>(gemm_node);
6✔
564
}
6✔
565

566
GEMMNode& add_gemm_node(
567
    builder::StructuredSDFGBuilder& builder,
568
    Block& block,
569
    const std::string& ptr_a,
570
    const std::string& ptr_b,
571
    const std::string& ptr_c,
572
    data_flow::AccessNode& alpha_node,
573
    data_flow::AccessNode& beta_node,
574
    const BLAS_Precision& precision,
575
    const BLAS_Layout& layout,
576
    const BLAS_Transpose& trans_a,
577
    const BLAS_Transpose& trans_b,
578
    symbolic::Expression& m,
579
    symbolic::Expression& n,
580
    symbolic::Expression& k,
581
    symbolic::Expression& lda,
582
    symbolic::Expression& ldb,
583
    symbolic::Expression& ldc,
584
    const types::IType& ptr_type,
585
    const types::IType& factor_type,
586
    DebugInfo debug_info,
587
    data_flow::ImplementationType impl_type
588
) {
×
589
    return add_gemm_node(
×
590
        builder,
×
591
        block,
×
592
        ptr_a,
×
593
        ptr_b,
×
594
        ptr_c,
×
595
        alpha_node,
×
596
        beta_node,
×
597
        precision,
×
598
        layout,
×
599
        trans_a,
×
600
        trans_b,
×
601
        m,
×
602
        n,
×
603
        k,
×
604
        lda,
×
605
        ldb,
×
606
        ldc,
×
607
        ptr_type,
×
608
        ptr_type,
×
609
        ptr_type,
×
610
        factor_type,
×
611
        debug_info,
×
612
        debug_info,
×
613
        debug_info,
×
614
        debug_info,
×
615
        debug_info,
×
616
        debug_info,
×
617
        debug_info,
×
618
        impl_type
×
619
    );
×
620
}
×
621

622
} // namespace blas
623
} // namespace math
624
} // 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