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

daisytuner / docc / 26556322966

27 May 2026 03:45PM UTC coverage: 60.869% (-0.02%) from 60.886%
26556322966

push

github

web-flow
Libnode ptr edges (#719)

Migrating SDFGs to treat pointers as inputs to libNodes / Calls as scalars.
A pointer will only appear in an output edge if its actually returned from the function (like malloc).

* Stdlib, Blas and Tensor Matmul nodes were migrated to this new format. Other, currently transitory Tensor Nodes are not yet migrated.
* DOCC version was bumped to incorporate previous docc-llvm versions (up to 0.4.0) that had been counted separately.
! Until all passes consider the use / leak of pointers as uncertainty / hiding potential writes, TensorNodes are declared as general side-effect.
* Lots of utility functions to centralize the creation (and edges) of various libNodes that needed to be changed.
* Fixed & unified docc paths across python and llvm front-ends.
* Skip BlockFusion test that fails to its libNodes currently having side effects
~ Prevent a crash in DotViz when using symbolic offsets into structs
* Removing old ConstProp pass, it is not safe for the new pointer representation and should not be all too critical

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

87 existing lines in 28 files now uncovered.

35225 of 57870 relevant lines covered (60.87%)

11046.32 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