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

daisytuner / docc / 26520678771

27 May 2026 03:22PM UTC coverage: 60.864% (-0.02%) from 60.886%
26520678771

Pull #719

github

web-flow
Merge 99c5e4f9d into 707dadcf8
Pull Request #719: Libnode ptr edges

961 of 1749 new or added lines in 52 files covered. (54.95%)

90 existing lines in 29 files now uncovered.

35222 of 57870 relevant lines covered (60.86%)

11043.61 hits per line

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

43.58
/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
#include "sdfg/analysis/scope_analysis.h"
7

8
namespace sdfg {
9
namespace math {
10
namespace blas {
11

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

42
BLAS_Layout GEMMNode::layout() const { return this->layout_; };
3✔
43

44
BLAS_Transpose GEMMNode::trans_a() const { return this->trans_a_; };
3✔
45

46
BLAS_Transpose GEMMNode::trans_b() const { return this->trans_b_; };
3✔
47

48
symbolic::Expression GEMMNode::m() const { return this->m_; };
14✔
49

50
symbolic::Expression GEMMNode::n() const { return this->n_; };
14✔
51

52
symbolic::Expression GEMMNode::k() const { return this->k_; };
14✔
53

54
symbolic::Expression GEMMNode::lda() const { return this->lda_; };
8✔
55

56
symbolic::Expression GEMMNode::ldb() const { return this->ldb_; };
8✔
57

58
symbolic::Expression GEMMNode::ldc() const { return this->ldc_; };
8✔
59

60
symbolic::SymbolSet GEMMNode::symbols() const {
×
61
    symbolic::SymbolSet syms;
×
62

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

82
    return syms;
×
83
};
×
84

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

94
void GEMMNode::validate(const Function& function) const { BLASNode::validate(function); }
8✔
95

96
bool GEMMNode::expand(builder::StructuredSDFGBuilder& builder, analysis::AnalysisManager& analysis_manager) {
6✔
97
    auto& scope_analysis = analysis_manager.get<analysis::ScopeAnalysis>();
6✔
98

99
    auto& dataflow = this->get_parent();
6✔
100
    auto& block = static_cast<structured_control_flow::Block&>(*dataflow.get_parent());
6✔
101
    auto& parent = static_cast<structured_control_flow::Sequence&>(*scope_analysis.parent_scope(&block));
6✔
102
    int index = parent.index(block);
6✔
103
    auto& transition = parent.at(index).second;
6✔
104

105
    if (trans_a_ == BLAS_Transpose::ConjTrans || trans_b_ == BLAS_Transpose::ConjTrans) {
6✔
106
        return false;
×
107
    }
×
108

109
    auto primitive_type = scalar_primitive();
6✔
110
    if (primitive_type == types::PrimitiveType::Void) {
6✔
111
        return false;
×
112
    }
×
113

114
    types::Scalar scalar_type(primitive_type);
6✔
115

116
    auto in_edges = dataflow.in_edges(*this);
6✔
117
    auto in_edges_it = in_edges.begin();
6✔
118

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

143
    // Checks if legal
144
    auto* input_node_a = static_cast<data_flow::AccessNode*>(&iedge_a->src());
6✔
145
    auto* input_node_b = static_cast<data_flow::AccessNode*>(&iedge_b->src());
6✔
146
    auto* input_node_c = static_cast<data_flow::AccessNode*>(&iedge_c->src());
6✔
147
    auto* alpha_node = static_cast<data_flow::AccessNode*>(&alpha_edge->src());
6✔
148
    auto* beta_node = static_cast<data_flow::AccessNode*>(&beta_edge->src());
6✔
149

150
    // we must be the only thing in this block, as we do not support splitting a block into pre, expanded lib-node, post
151
    if (!input_node_a || dataflow.in_degree(*input_node_a) != 0 || !input_node_b ||
6✔
152
        dataflow.in_degree(*input_node_b) != 0 || !input_node_c || dataflow.in_degree(*input_node_c) != 0) {
6✔
UNCOV
153
        return false; // data nodes are not standalone
×
154
    }
×
155
    if (dataflow.in_degree(*alpha_node) != 0 || dataflow.in_degree(*beta_node) != 0) {
6✔
156
        return false; // alpha and beta are not standalone
×
157
    }
×
158
    for (auto* nd : dataflow.data_nodes()) {
30✔
159
        if (nd != input_node_a && nd != input_node_b && nd != input_node_c && (!alpha_node || nd != alpha_node) &&
30✔
160
            (!beta_node || nd != beta_node)) {
30✔
161
            return false; // there are other nodes in here that we could not preserve correctly
×
162
        }
×
163
    }
30✔
164

165
    auto& A_var = input_node_a->data();
6✔
166
    auto& B_var = input_node_b->data();
6✔
167
    auto& C_ptr = input_node_c->data();
6✔
168

169

170
    // Add new graph after the current block
171
    auto& new_sequence = builder.add_sequence_before(parent, block, transition.assignments(), block.debug_info());
6✔
172

173
    // Add maps
174
    std::vector<symbolic::Expression> indvar_ends{this->m(), this->n(), this->k()};
6✔
175
    data_flow::Subset new_subset;
6✔
176
    structured_control_flow::Sequence* last_scope = &new_sequence;
6✔
177
    structured_control_flow::StructuredLoop* last_map = nullptr;
6✔
178
    structured_control_flow::StructuredLoop* output_loop = nullptr;
6✔
179
    std::vector<std::string> indvar_names{"_i", "_j", "_k"};
6✔
180

181
    std::string sum_var = builder.find_new_name("_sum");
6✔
182
    builder.add_container(sum_var, scalar_type);
6✔
183

184
    for (size_t i = 0; i < 3; i++) {
24✔
185
        auto dim_begin = symbolic::zero();
18✔
186
        auto& dim_end = indvar_ends[i];
18✔
187

188
        std::string indvar_str = builder.find_new_name(indvar_names[i]);
18✔
189
        builder.add_container(indvar_str, types::Scalar(types::PrimitiveType::UInt64));
18✔
190

191
        auto indvar = symbolic::symbol(indvar_str);
18✔
192
        auto init = dim_begin;
18✔
193
        auto update = symbolic::add(indvar, symbolic::one());
18✔
194
        auto condition = symbolic::Lt(indvar, dim_end);
18✔
195
        if (i < 2) {
18✔
196
            last_map = &builder.add_map(
12✔
197
                *last_scope,
12✔
198
                indvar,
12✔
199
                condition,
12✔
200
                init,
12✔
201
                update,
12✔
202
                structured_control_flow::ScheduleType_Sequential::create(),
12✔
203
                {},
12✔
204
                block.debug_info()
12✔
205
            );
12✔
206
        } else {
12✔
207
            last_map = &builder.add_for(*last_scope, indvar, condition, init, update, {}, block.debug_info());
6✔
208
        }
6✔
209
        last_scope = &last_map->root();
18✔
210

211
        if (i == 1) {
18✔
212
            output_loop = last_map;
6✔
213
        }
6✔
214

215
        new_subset.push_back(indvar);
18✔
216
    }
18✔
217

218

219
    // Add code
220
    auto& init_block = builder.add_block_before(output_loop->root(), *last_map, {}, block.debug_info());
6✔
221
    auto& sum_init = builder.add_access(init_block, sum_var, block.debug_info());
6✔
222

223
    auto& zero_node = builder.add_constant(init_block, "0.0", alpha_edge->base_type(), block.debug_info());
6✔
224
    auto& init_tasklet = builder.add_tasklet(init_block, data_flow::assign, "_out", {"_in"}, block.debug_info());
6✔
225
    builder.add_computational_memlet(init_block, zero_node, init_tasklet, "_in", {}, block.debug_info());
6✔
226
    builder.add_computational_memlet(init_block, init_tasklet, "_out", sum_init, {}, block.debug_info());
6✔
227

228
    auto& code_block = builder.add_block(*last_scope, {}, block.debug_info());
6✔
229
    auto& input_node_a_new = builder.add_access(code_block, A_var, input_node_a->debug_info());
6✔
230
    auto& input_node_b_new = builder.add_access(code_block, B_var, input_node_b->debug_info());
6✔
231

232
    auto& core_fma =
6✔
233
        builder.add_tasklet(code_block, data_flow::fp_fma, "_out", {"_in1", "_in2", "_in3"}, block.debug_info());
6✔
234
    auto& sum_in = builder.add_access(code_block, sum_var, block.debug_info());
6✔
235
    auto& sum_out = builder.add_access(code_block, sum_var, block.debug_info());
6✔
236
    builder.add_computational_memlet(code_block, sum_in, core_fma, "_in3", {}, block.debug_info());
6✔
237

238
    // Row-major indexing: address = ld * row + col
239
    // No transpose: A is m×k, access A[i, k] => lda*i + k
240
    // Transpose:    A is k×m stored, access A[k, i] => lda*k + i
241
    symbolic::Expression a_idx = (trans_a_ == BLAS_Transpose::Trans)
6✔
242
                                     ? symbolic::add(symbolic::mul(lda(), new_subset[2]), new_subset[0])
6✔
243
                                     : symbolic::add(symbolic::mul(lda(), new_subset[0]), new_subset[2]);
6✔
244
    builder.add_computational_memlet(
6✔
245
        code_block, input_node_a_new, core_fma, "_in1", {a_idx}, iedge_a->base_type(), iedge_a->debug_info()
6✔
246
    );
6✔
247
    // No transpose: B is k×n, access B[k, j] => ldb*k + j
248
    // Transpose:    B is n×k stored, access B[j, k] => ldb*j + k
249
    symbolic::Expression b_idx = (trans_b_ == BLAS_Transpose::Trans)
6✔
250
                                     ? symbolic::add(symbolic::mul(ldb(), new_subset[1]), new_subset[2])
6✔
251
                                     : symbolic::add(symbolic::mul(ldb(), new_subset[2]), new_subset[1]);
6✔
252
    builder.add_computational_memlet(
6✔
253
        code_block, input_node_b_new, core_fma, "_in2", {b_idx}, iedge_b->base_type(), iedge_b->debug_info()
6✔
254
    );
6✔
255
    builder.add_computational_memlet(code_block, core_fma, "_out", sum_out, {}, iedge_c->debug_info());
6✔
256

257
    auto& flush_block = builder.add_block_after(output_loop->root(), *last_map, {}, block.debug_info());
6✔
258
    auto& sum_final = builder.add_access(flush_block, sum_var, block.debug_info());
6✔
259
    auto& input_node_c_new = builder.add_access(flush_block, C_ptr, input_node_c->debug_info());
6✔
260
    symbolic::Expression c_idx = symbolic::add(symbolic::mul(ldc(), new_subset[0]), new_subset[1]);
6✔
261

262
    auto& scale_sum_tasklet =
6✔
263
        builder.add_tasklet(flush_block, data_flow::TaskletCode::fp_mul, "_out", {"_in1", "_in2"}, block.debug_info());
6✔
264
    builder.add_computational_memlet(flush_block, sum_final, scale_sum_tasklet, "_in1", {}, block.debug_info());
6✔
265
    if (auto const_node = dynamic_cast<data_flow::ConstantNode*>(alpha_node)) {
6✔
266
        auto& alpha_node_new =
6✔
267
            builder.add_constant(flush_block, const_node->data(), const_node->type(), block.debug_info());
6✔
268
        builder.add_computational_memlet(flush_block, alpha_node_new, scale_sum_tasklet, "_in2", {}, block.debug_info());
6✔
269
    } else {
6✔
270
        auto& alpha_node_new = builder.add_access(flush_block, alpha_node->data(), block.debug_info());
×
271
        builder.add_computational_memlet(flush_block, alpha_node_new, scale_sum_tasklet, "_in2", {}, block.debug_info());
×
272
    }
×
273

274
    std::string scaled_sum_temp = builder.find_new_name("scaled_sum_temp");
6✔
275
    builder.add_container(scaled_sum_temp, scalar_type);
6✔
276
    auto& scaled_sum_final = builder.add_access(flush_block, scaled_sum_temp, block.debug_info());
6✔
277
    builder.add_computational_memlet(
6✔
278
        flush_block, scale_sum_tasklet, "_out", scaled_sum_final, {}, scalar_type, block.debug_info()
6✔
279
    );
6✔
280

281
    auto& scale_input_tasklet =
6✔
282
        builder.add_tasklet(flush_block, data_flow::TaskletCode::fp_mul, "_out", {"_in1", "_in2"}, block.debug_info());
6✔
283
    builder.add_computational_memlet(
6✔
284
        flush_block, input_node_c_new, scale_input_tasklet, "_in1", {c_idx}, iedge_c->base_type(), iedge_c->debug_info()
6✔
285
    );
6✔
286
    if (auto const_node = dynamic_cast<data_flow::ConstantNode*>(beta_node)) {
6✔
287
        auto& beta_node_new =
6✔
288
            builder.add_constant(flush_block, const_node->data(), const_node->type(), block.debug_info());
6✔
289
        builder
6✔
290
            .add_computational_memlet(flush_block, beta_node_new, scale_input_tasklet, "_in2", {}, block.debug_info());
6✔
291
    } else {
6✔
292
        auto& beta_node_new = builder.add_access(flush_block, beta_node->data(), block.debug_info());
×
293
        builder
×
294
            .add_computational_memlet(flush_block, beta_node_new, scale_input_tasklet, "_in2", {}, block.debug_info());
×
295
    }
×
296

297
    std::string scaled_input_temp = builder.find_new_name("scaled_input_temp");
6✔
298
    builder.add_container(scaled_input_temp, scalar_type);
6✔
299
    auto& scaled_input_c = builder.add_access(flush_block, scaled_input_temp, block.debug_info());
6✔
300
    builder.add_computational_memlet(
6✔
301
        flush_block, scale_input_tasklet, "_out", scaled_input_c, {}, scalar_type, block.debug_info()
6✔
302
    );
6✔
303

304
    auto& flush_add_tasklet =
6✔
305
        builder.add_tasklet(flush_block, data_flow::TaskletCode::fp_add, "_out", {"_in1", "_in2"}, block.debug_info());
6✔
306
    auto& output_node_new = builder.add_access(flush_block, C_ptr, input_node_c->debug_info());
6✔
307
    builder.add_computational_memlet(
6✔
308
        flush_block, scaled_sum_final, flush_add_tasklet, "_in1", {}, scalar_type, block.debug_info()
6✔
309
    );
6✔
310
    builder.add_computational_memlet(
6✔
311
        flush_block, scaled_input_c, flush_add_tasklet, "_in2", {}, scalar_type, block.debug_info()
6✔
312
    );
6✔
313
    builder.add_computational_memlet(
6✔
314
        flush_block, flush_add_tasklet, "_out", output_node_new, {c_idx}, iedge_c->base_type(), iedge_c->debug_info()
6✔
315
    );
6✔
316

317

318
    // Clean up block
319
    builder.remove_memlet(block, *iedge_a);
6✔
320
    builder.remove_memlet(block, *iedge_b);
6✔
321
    builder.remove_memlet(block, *iedge_c);
6✔
322
    builder.remove_memlet(block, *alpha_edge);
6✔
323
    builder.remove_node(block, *alpha_node);
6✔
324
    builder.remove_memlet(block, *beta_edge);
6✔
325
    builder.remove_node(block, *beta_node);
6✔
326
    builder.remove_node(block, *input_node_a);
6✔
327
    builder.remove_node(block, *input_node_b);
6✔
328
    builder.remove_node(block, *input_node_c);
6✔
329
    builder.remove_node(block, *this);
6✔
330
    builder.remove_child(parent, index + 1);
6✔
331

332
    return true;
6✔
333
}
6✔
334

335
symbolic::Expression GEMMNode::flop() const {
×
336
    return flops(symbolic::__true__(), symbolic::__true__(), symbolic::__true__(), symbolic::__true__());
×
337
}
×
338

339
symbolic::Expression GEMMNode::flops(
340
    symbolic::Condition alpha_non_zero,
341
    symbolic::Condition alpha_non_ident,
342
    symbolic::Condition beta_non_zero,
343
    symbolic::Condition beta_non_ident
344
) const {
×
345
    auto res_elems = symbolic::mul(this->m_, this->n_);
×
346

347
    // conditional on alpha != 0.0
348
    auto mm_mul_ops = symbolic::mul(symbolic::mul(res_elems, this->k_), alpha_non_zero);
×
349
    auto mm_sum_ops = symbolic::mul(symbolic::mul(res_elems, symbolic::sub(this->k_, symbolic::one())), alpha_non_zero);
×
350
    // conditional on alpha != 1.0 && alpha != 0.0
351
    auto mm_alpha_scale_ops = symbolic::mul(res_elems, symbolic::And(alpha_non_ident, alpha_non_zero));
×
352
    // conditional on beta != 1.0 && beta != 0.0
353
    auto mm_beta_scale_ops = symbolic::mul(res_elems, symbolic::And(beta_non_ident, beta_non_zero));
×
354
    auto mm_beta_scaled_sum_ops = symbolic::mul(res_elems, beta_non_zero);
×
355
    auto mul_ops = symbolic::add(mm_mul_ops, symbolic::add(mm_alpha_scale_ops, mm_beta_scale_ops));
×
356
    auto add_ops = symbolic::add(mm_sum_ops, mm_beta_scaled_sum_ops);
×
357
    return symbolic::add(mul_ops, add_ops);
×
358
}
×
359

360
std::unique_ptr<data_flow::DataFlowNode> GEMMNode::
361
    clone(size_t element_id, const graph::Vertex vertex, data_flow::DataFlowGraph& parent) const {
×
362
    auto node_clone = std::unique_ptr<GEMMNode>(new GEMMNode(
×
363
        element_id,
×
364
        this->debug_info(),
×
365
        vertex,
×
366
        parent,
×
367
        this->implementation_type_,
×
368
        this->precision_,
×
369
        this->layout_,
×
370
        this->trans_a_,
×
371
        this->trans_b_,
×
372
        this->m_,
×
373
        this->n_,
×
374
        this->k_,
×
375
        this->lda_,
×
376
        this->ldb_,
×
377
        this->ldc_
×
378
    ));
×
379
    return std::move(node_clone);
×
380
}
×
381

382
std::string GEMMNode::toStr() const {
×
383
    return LibraryNode::toStr() + "(" + static_cast<char>(precision_) + ", " +
×
384
           std::string(BLAS_Layout_to_short_string(layout_)) + ", " + BLAS_Transpose_to_char(trans_a_) +
×
385
           BLAS_Transpose_to_char(trans_b_) + ", " + m_->__str__() + ", " + n_->__str__() + ", " + k_->__str__() +
×
386
           ", " + lda_->__str__() + ", " + ldb_->__str__() + ", " + ldc_->__str__() + ")";
×
387
}
×
388

389
symbolic::Expression GEMMNode::calc_matrix_access_range(
390
    const symbolic::Expression& outer_dim,
391
    const symbolic::Expression& inner_dim,
392
    const symbolic::Expression& line_size,
393
    BLAS_Transpose trans,
394
    BLAS_Layout layout
NEW
395
) {
×
NEW
396
    if ((trans == BLAS_Transpose::No) ^ (layout == BLAS_Layout::ColMajor)) {
×
NEW
397
        return symbolic::mul(outer_dim, line_size);
×
NEW
398
    } else {
×
NEW
399
        return symbolic::mul(inner_dim, line_size);
×
NEW
400
    }
×
NEW
401
}
×
402

403

NEW
404
data_flow::PointerAccessType GEMMNode::pointer_access_type(int input_idx) const {
×
NEW
405
    if (input_idx == 0) { // A: m x k
×
NEW
406
        return data_flow::PointerAccessMeta::
×
NEW
407
            create_read_only(calc_matrix_access_range(m_, k_, lda_, trans_a_, layout_), true);
×
NEW
408
    } else if (input_idx == 1) { // B: k x n
×
NEW
409
        return data_flow::PointerAccessMeta::
×
NEW
410
            create_read_only(calc_matrix_access_range(k_, n_, ldb_, trans_b_, layout_), true);
×
NEW
411
    } else if (input_idx == 2) {
×
412
        // for beta == 0, there would no reads of C. But we currently have no mechanism to access const-prop knowledge
413
        // like tha
NEW
414
        if (symbolic::eq(ldc_, n_)) { // non-sparse access over the m x n range
×
NEW
415
            return data_flow::PointerAccessMeta::
×
NEW
416
                create_full_write_only(calc_matrix_access_range(m_, n_, ldc_, BLAS_Transpose::No, layout_), true);
×
NEW
417
        } else {
×
418
            // sparse access. But with only Convex Pattern for now, we cannot represent which values are
NEW
419
            auto pattern =
×
NEW
420
                data_flow::ConvexAccessPattern::create(calc_matrix_access_range(m_, n_, ldc_, BLAS_Transpose::No, layout_)
×
NEW
421
                );
×
422
            // full-overwritten and which are DC.
NEW
423
            return data_flow::PointerAccessMeta::create_generic(pattern->ref(), std::move(pattern), true);
×
NEW
424
        }
×
NEW
425
    } else {
×
NEW
426
        return LibraryNode::pointer_access_type(input_idx);
×
NEW
427
    }
×
NEW
428
}
×
429

430
nlohmann::json GEMMNodeSerializer::serialize(const data_flow::LibraryNode& library_node) {
×
431
    const GEMMNode& gemm_node = static_cast<const GEMMNode&>(library_node);
×
432
    nlohmann::json j;
×
433

434
    serializer::JSONSerializer serializer;
×
435
    j["code"] = gemm_node.code().value();
×
436
    j["precision"] = gemm_node.precision();
×
437
    j["layout"] = gemm_node.layout();
×
438
    j["trans_a"] = gemm_node.trans_a();
×
439
    j["trans_b"] = gemm_node.trans_b();
×
440
    j["m"] = serializer.expression(gemm_node.m());
×
441
    j["n"] = serializer.expression(gemm_node.n());
×
442
    j["k"] = serializer.expression(gemm_node.k());
×
443
    j["lda"] = serializer.expression(gemm_node.lda());
×
444
    j["ldb"] = serializer.expression(gemm_node.ldb());
×
445
    j["ldc"] = serializer.expression(gemm_node.ldc());
×
446

447
    return j;
×
448
}
×
449

450
data_flow::LibraryNode& GEMMNodeSerializer::deserialize(
451
    const nlohmann::json& j, builder::StructuredSDFGBuilder& builder, structured_control_flow::Block& parent
452
) {
×
453
    // Assertions for required fields
454
    assert(j.contains("element_id"));
×
455
    assert(j.contains("code"));
×
456
    assert(j.contains("debug_info"));
×
457

458
    auto code = j["code"].get<std::string>();
×
459
    if (code != LibraryNodeType_GEMM.value()) {
×
460
        throw std::runtime_error("Invalid library node code");
×
461
    }
×
462

463
    // Extract debug info using JSONSerializer
464
    sdfg::serializer::JSONSerializer serializer;
×
465
    DebugInfo debug_info = serializer.json_to_debug_info(j["debug_info"]);
×
466

467
    auto precision = j.at("precision").get<BLAS_Precision>();
×
468
    auto layout = j.at("layout").get<BLAS_Layout>();
×
469
    auto trans_a = j.at("trans_a").get<BLAS_Transpose>();
×
470
    auto trans_b = j.at("trans_b").get<BLAS_Transpose>();
×
471
    auto m = symbolic::parse(j.at("m"));
×
472
    auto n = symbolic::parse(j.at("n"));
×
473
    auto k = symbolic::parse(j.at("k"));
×
474
    auto lda = symbolic::parse(j.at("lda"));
×
475
    auto ldb = symbolic::parse(j.at("ldb"));
×
476
    auto ldc = symbolic::parse(j.at("ldc"));
×
477

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

480
    return builder.add_library_node<
×
481
        GEMMNode>(parent, debug_info, implementation_type, precision, layout, trans_a, trans_b, m, n, k, lda, ldb, ldc);
×
482
}
×
483

484
GEMMNodeDispatcher_BLAS::GEMMNodeDispatcher_BLAS(
485
    codegen::LanguageExtension& language_extension,
486
    const Function& function,
487
    const data_flow::DataFlowGraph& data_flow_graph,
488
    const GEMMNode& node
489
)
490
    : codegen::LibraryNodeDispatcher(language_extension, function, data_flow_graph, node) {}
×
491

492
void GEMMNodeDispatcher_BLAS::dispatch_code_with_edges(
493
    codegen::CodegenOutput& out,
494
    std::vector<codegen::DispatchInput>& inputs,
495
    std::vector<codegen::DispatchOutput>& outputs
496
) {
×
497
    auto& gemm_node = static_cast<const GEMMNode&>(this->node_);
×
498

499
    sdfg::types::Scalar base_type(types::PrimitiveType::Void);
×
500
    switch (gemm_node.precision()) {
×
501
        case BLAS_Precision::h:
×
502
            base_type = types::Scalar(types::PrimitiveType::Half);
×
503
            break;
×
504
        case BLAS_Precision::s:
×
505
            base_type = types::Scalar(types::PrimitiveType::Float);
×
506
            break;
×
507
        case BLAS_Precision::d:
×
508
            base_type = types::Scalar(types::PrimitiveType::Double);
×
509
            break;
×
510
        default:
×
511
            throw std::runtime_error("Invalid BLAS_Precision value");
×
512
    }
×
513

NEW
514
    out.library_snippet_factory.require_dependency(BLASLibDependency::instance());
×
515

NEW
516
    out.stream << "cblas_" << BLAS_Precision_to_string(gemm_node.precision()) << "gemm(";
×
NEW
517
    out.stream.changeIndent(+4);
×
NEW
518
    out.stream << BLAS_Layout_to_string(gemm_node.layout());
×
NEW
519
    out.stream << ", ";
×
NEW
520
    out.stream << BLAS_Transpose_to_string(gemm_node.trans_a());
×
NEW
521
    out.stream << ", ";
×
NEW
522
    out.stream << BLAS_Transpose_to_string(gemm_node.trans_b());
×
NEW
523
    out.stream << ", ";
×
NEW
524
    out.stream << this->language_extension_.expression(gemm_node.m());
×
NEW
525
    out.stream << ", ";
×
NEW
526
    out.stream << this->language_extension_.expression(gemm_node.n());
×
NEW
527
    out.stream << ", ";
×
NEW
528
    out.stream << this->language_extension_.expression(gemm_node.k());
×
NEW
529
    out.stream << ", ";
×
NEW
530
    out.stream << inputs.at(GEMMNode::ALPHA_INPUT_IDX).expr;
×
NEW
531
    out.stream << ", ";
×
NEW
532
    out.stream << inputs.at(GEMMNode::A_INPUT_IDX).expr;
×
NEW
533
    out.stream << ", ";
×
NEW
534
    out.stream << this->language_extension_.expression(gemm_node.lda());
×
NEW
535
    out.stream << ", ";
×
NEW
536
    out.stream << inputs.at(GEMMNode::B_INPUT_IDX).expr;
×
NEW
537
    out.stream << ", ";
×
NEW
538
    out.stream << this->language_extension_.expression(gemm_node.ldb());
×
NEW
539
    out.stream << ", ";
×
NEW
540
    out.stream << inputs.at(GEMMNode::BETA_INPUT_IDX).expr;
×
NEW
541
    out.stream << ", ";
×
NEW
542
    out.stream << inputs.at(GEMMNode::C_INPUT_IDX).expr;
×
NEW
543
    out.stream << ", ";
×
NEW
544
    out.stream << this->language_extension_.expression(gemm_node.ldc());
×
545

NEW
546
    out.stream.changeIndent(-4);
×
NEW
547
    out.stream << ");" << std::endl;
×
NEW
548
}
×
549

550
GEMMNode& add_gemm_node(
551
    builder::StructuredSDFGBuilder& builder,
552
    Block& block,
553
    const std::string& ptr_a,
554
    const std::string& ptr_b,
555
    const std::string& ptr_c,
556
    data_flow::AccessNode& alpha_node,
557
    data_flow::AccessNode& beta_node,
558
    const BLAS_Precision& precision,
559
    const BLAS_Layout& layout,
560
    const BLAS_Transpose& trans_a,
561
    const BLAS_Transpose& trans_b,
562
    symbolic::Expression& m,
563
    symbolic::Expression& n,
564
    symbolic::Expression& k,
565
    symbolic::Expression& lda,
566
    symbolic::Expression& ldb,
567
    symbolic::Expression& ldc,
568
    const types::IType& a_type,
569
    const types::IType& b_type,
570
    const types::IType& c_type,
571
    const types::IType& factor_type,
572
    DebugInfo debug_info,
573
    DebugInfo a_access_deb_info,
574
    DebugInfo b_access_deb_info,
575
    DebugInfo c_access_deb_info,
576
    DebugInfo a_edge_deb_info,
577
    DebugInfo b_edge_deb_info,
578
    DebugInfo c_edge_deb_info,
579
    data_flow::ImplementationType impl_type
NEW
580
) {
×
NEW
581
    auto& gemm_node = builder.add_library_node<sdfg::math::blas::GEMMNode>(
×
NEW
582
        block, debug_info, std::move(impl_type), precision, layout, trans_a, trans_b, m, n, k, lda, ldb, ldc
×
NEW
583
    );
×
584

585
    // Add access nodes
NEW
586
    auto& a_node_in = builder.add_access(block, ptr_a, a_access_deb_info);
×
NEW
587
    auto& b_node_in = builder.add_access(block, ptr_b, b_access_deb_info);
×
NEW
588
    auto& c_node_in = builder.add_access(block, ptr_c, c_access_deb_info);
×
589

590
    // Add edges
NEW
591
    builder.add_computational_memlet(block, a_node_in, gemm_node, "__A", {}, a_type, a_edge_deb_info);
×
NEW
592
    builder.add_computational_memlet(block, b_node_in, gemm_node, "__B", {}, b_type, b_edge_deb_info);
×
NEW
593
    builder.add_computational_memlet(block, c_node_in, gemm_node, "__C", {}, c_type, c_edge_deb_info);
×
NEW
594
    builder.add_computational_memlet(block, alpha_node, gemm_node, "__alpha", {}, factor_type, debug_info);
×
NEW
595
    builder.add_computational_memlet(block, beta_node, gemm_node, "__beta", {}, factor_type, debug_info);
×
596

NEW
597
    return static_cast<GEMMNode&>(gemm_node);
×
NEW
598
}
×
599

600
GEMMNode& add_gemm_node(
601
    builder::StructuredSDFGBuilder& builder,
602
    Block& block,
603
    const std::string& ptr_a,
604
    const std::string& ptr_b,
605
    const std::string& ptr_c,
606
    data_flow::AccessNode& alpha_node,
607
    data_flow::AccessNode& beta_node,
608
    const BLAS_Precision& precision,
609
    const BLAS_Layout& layout,
610
    const BLAS_Transpose& trans_a,
611
    const BLAS_Transpose& trans_b,
612
    symbolic::Expression& m,
613
    symbolic::Expression& n,
614
    symbolic::Expression& k,
615
    symbolic::Expression& lda,
616
    symbolic::Expression& ldb,
617
    symbolic::Expression& ldc,
618
    const types::IType& ptr_type,
619
    const types::IType& factor_type,
620
    DebugInfo debug_info,
621
    data_flow::ImplementationType impl_type
NEW
622
) {
×
NEW
623
    return add_gemm_node(
×
NEW
624
        builder,
×
NEW
625
        block,
×
NEW
626
        ptr_a,
×
NEW
627
        ptr_b,
×
NEW
628
        ptr_c,
×
NEW
629
        alpha_node,
×
NEW
630
        beta_node,
×
NEW
631
        precision,
×
NEW
632
        layout,
×
NEW
633
        trans_a,
×
NEW
634
        trans_b,
×
NEW
635
        m,
×
NEW
636
        n,
×
NEW
637
        k,
×
NEW
638
        lda,
×
NEW
639
        ldb,
×
NEW
640
        ldc,
×
NEW
641
        ptr_type,
×
NEW
642
        ptr_type,
×
NEW
643
        ptr_type,
×
NEW
644
        factor_type,
×
NEW
645
        debug_info,
×
NEW
646
        debug_info,
×
NEW
647
        debug_info,
×
NEW
648
        debug_info,
×
NEW
649
        debug_info,
×
NEW
650
        debug_info,
×
NEW
651
        debug_info,
×
NEW
652
        impl_type
×
NEW
653
    );
×
UNCOV
654
}
×
655

656
} // namespace blas
657
} // namespace math
658
} // 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