Repository navigation
Expand file tree
/
Copy pathconv.cpp
More file actions
92 lines (77 loc) · 3.08 KB
/
Copy pathconv.cpp
File metadata and controls
92 lines (77 loc) · 3.08 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
/* Copyright (c) 2021-2025 Advanced Micro Devices, Inc. All rights reserved. */
#include "shaders/operators/conv.hpp"
#include "impl/conv/utils.hpp"
#include "impl/conv/mxn/convMxN.hpp"
#include "impl/conv/1x1/conv1x1.hpp"
#include "impl/conv/dilated/dilatedConv.hpp"
#include "impl/conv/depthWise/depthWiseConv.hpp"
template class mlss::CaseBase<mlss::conv::mxn::ConvMxN, mlss::op::OperatorConv>;
template class mlss::CaseBase<mlss::conv::one_by_one::Conv1x1, mlss::op::OperatorConv>;
template class mlss::CaseBase<mlss::conv::dilated::DilatedConv, mlss::op::OperatorConv>;
template class mlss::CaseBase<mlss::conv::depth_wise::DepthWiseConv, mlss::op::OperatorConv>;
namespace mlss::op
{
OperatorConv::OperatorConv(const std::vector<mlss::Attribute>& attributes, GfxIpTriple gfxip)
: base(attributes, gfxip)
{
this->m_implName = "Conv";
}
std::string OperatorConv::getOperatorName()
{
return "AMDMLSS::OperatorConv";
}
std::expected<Binaries, std::error_code> OperatorConv::getBinaries() const
{
auto* conv1x1 = CaseRegistry<OperatorConv>::get<mlss::conv::one_by_one::Conv1x1>();
if (conv1x1 != nullptr)
{
auto result = conv1x1->getBinaries(m_attributes, m_gfxIpTriple);
if (result.has_value())
{
m_implName = conv1x1->name;
return result;
}
}
auto* convMxN = CaseRegistry<OperatorConv>::get<mlss::conv::mxn::ConvMxN>();
if (convMxN != nullptr)
{
auto result = convMxN->getBinaries(m_attributes, m_gfxIpTriple);
if (result.has_value())
{
m_implName = convMxN->name;
return result;
}
}
auto* dilated = CaseRegistry<OperatorConv>::get<mlss::conv::dilated::DilatedConv>();
if (dilated != nullptr)
{
auto result = dilated->getBinaries(m_attributes, m_gfxIpTriple);
if (result.has_value())
{
m_implName = dilated->name;
return result;
}
}
return std::unexpected(std::make_error_code(std::errc::not_supported));
}
bool OperatorConv::getCapsImpl(const std::vector<mlss::Attribute>& attributes, GfxIpTriple gfxip)
{
auto params = mlss::conv::utils::buildConvParams(attributes);
auto* conv1x1 = CaseRegistry<OperatorConv>::get<mlss::conv::one_by_one::Conv1x1>();
if (conv1x1 != nullptr && conv1x1->getCaps(attributes, gfxip, ¶ms) != 0x00000000u)
{
return true;
}
auto* convMxN = CaseRegistry<OperatorConv>::get<mlss::conv::mxn::ConvMxN>();
if (convMxN != nullptr && convMxN->getCaps(attributes, gfxip, ¶ms) != 0x00000000u)
{
return true;
}
auto* dilated = CaseRegistry<OperatorConv>::get<mlss::conv::dilated::DilatedConv>();
if (dilated != nullptr && dilated->getCaps(attributes, gfxip, ¶ms) != 0x00000000u)
{
return true;
}
return false;
}
} // namespace mlss::op