Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ROperator_TopK.hxx
Go to the documentation of this file.
1#ifndef TMVA_SOFIE_ROPERATOR_TOPK
2#define TMVA_SOFIE_ROPERATOR_TOPK
3
5#include "TMVA/ROperator.hxx"
6#include "TMVA/RModel.hxx"
7
8#include <sstream>
9
10namespace TMVA {
11namespace Experimental {
12namespace SOFIE {
13
14template <typename T>
16
17private:
21
23 std::string fNK;
24 std::string fNX;
25 std::string fNVal;
26 std::string fNInd;
27 std::vector<Dim> fShapeX;
28 std::vector<Dim> fShapeY;
29 std::string fType;
30
31public:
33 ROperator_TopK(int attr_axis, int attr_largest, int attr_sorted, std::string nameK, std::string nameX, std::string nameVal, std::string nameInd)
37 fNK(UTILITY::Clean_name(nameK)),
38 fNX(UTILITY::Clean_name(nameX)),
39 fNVal(UTILITY::Clean_name(nameVal)),
40 fNInd(UTILITY::Clean_name(nameInd)){
43 }
44
45 void Initialize(RModel& model) override {
46 if (model.CheckIfTensorAlreadyExist(fNX) == false) {
47 // input must be a graph input, or already initialized intermediate tensor
48 throw std::runtime_error("TMVA SOFIE TopK Op Input Tensor is not found in model");
49 }
50 if (model.CheckIfTensorAlreadyExist(fNK) == false) {
51 // input must be a graph input, or already initialized intermediate tensor
52 throw std::runtime_error("TMVA SOFIE TopK Op Input Tensor i.e. K is not found in model");
53 }
54
55 fShapeX = model.GetDimTensorShape(fNX);
56 // K can either be an initialized tensor or a shape tensor, in which case its value is
57 // known only symbolically (e.g. it depends on one of the input dimensions)
58 Dim kdim;
59 if (model.IsShapeTensor(fNK)) {
60 auto &kvalues = model.GetShapeTensorValues(fNK);
61 if (kvalues.size() != 1)
62 throw std::runtime_error("TMVA SOFIE TopK Op input tensor K = " + fNK + " must be a single value");
63 kdim = kvalues[0];
64 } else if (model.IsInitializedTensor(fNK)) {
65 auto kptr = static_cast<int64_t *>(model.GetInitializedTensorData(fNK).get());
66 kdim = Dim{static_cast<size_t>(*kptr)};
67 model.SetNotWritableInitializedTensor(fNK);
68 } else {
69 throw std::runtime_error("TMVA SOFIE TopK Op input tensor K = " + fNK +
70 " must be known at initialization time");
71 }
73 if(static_cast<size_t>(fAttrAxis) >= fShapeX.size()){
74 throw
75 std::runtime_error("TMVA::SOFIE ONNX TopK op axis = "+ std::to_string(fAttrAxis) +" value exeeds size of tensor " +fNX+" of size "+ std::to_string(fShapeX.size()) +" .");
76 }
77 // fK cannot be larger that axis dimension
78 if (kdim.isParam || fShapeX[fAttrAxis].isParam)
79 fK = Dim{std::string("std::min(size_t(" + kdim.GetVal() + "), size_t(" + fShapeX[fAttrAxis].GetVal() + "))"),
80 static_cast<size_t>(-1)};
81 else
82 fK = Dim{std::min(kdim.dim, fShapeX[fAttrAxis].dim)};
83
84 // output shape is equal to input shape apart for value in fAttrAxis
87
88 model.AddIntermediateTensor(fNVal, model.GetTensorType(fNX), fShapeY);
89
90 // output indices should be an int64 tensor
91 model.AddIntermediateTensor(fNInd, ETensorType::INT64, fShapeY);
92 fType = ConvertTypeToString(model.GetTensorType(fNX));
93 model.AddNeededStdLib("algorithm");
94 model.AddNeededStdLib("cstdint");
95 model.AddNeededStdLib("cstring");
96
97 if (model.Verbose()) {
98 std::cout << "TopK " << fNX << " " << ConvertDimShapeToString(fShapeX)
99 << "---> " << fNVal << " " << ConvertDimShapeToString(fShapeY) << std::endl;
100 }
101 }
102
103 std::string Generate(std::string OpName) override {
104 OpName = "op_" + OpName;
105 if (fShapeX.empty()) {
106 throw std::runtime_error("TMVA SOFIE Operator TopK called to Generate without being initialized first");
107 }
108 std::stringstream out;
109 size_t size = fShapeX.size();
110 size_t axis = fAttrAxis < 0 ? size + fAttrAxis : fAttrAxis;
111 out << "\n" << SP << "//------ TopK\n";
112
116 // we perform loop on dimension before sorted axis and after sorted axis
117 std::vector<Dim> shape_before(fShapeX.begin(), fShapeX.begin() + axis); // input shape before axis
118 std::string n_before = (axis>0) ? ConvertDimShapeToLength(shape_before) : "1";
119 std::string n_after = strideX[axis].GetVal();
120 std::string n_elements = fShapeX[axis].GetVal(); // number of elements to be sorted
121
122 // }
123 out << SP << "{\n"; // to define a separate scope for the operator code
124
125 // Ties are broken by the element index, so no two entries ever compare equivalent:
126 // the ordering is total and the selected set is therefore unique. That is what makes
127 // the (unstable) std::nth_element below safe - it cannot pick a different set from a
128 // full sort.
129 //
130 // For float that ordering can be expressed as a single unsigned integer. Flipping the
131 // sign bit on positives and every bit on negatives maps a (non-NaN) float onto a
132 // uint32 whose unsigned order matches the float order; putting the element index in
133 // the low 32 bits then reproduces "ties by smaller index" exactly. A comparison
134 // becomes one 64-bit instruction instead of a two-field comparator call, and an
135 // element is 8 bytes instead of 16, which halves what nth_element has to move.
136 // Wider types cannot pack a value and an index into 64 bits, so they keep the pairs.
137 bool packed = (fType == "float");
138 // the index has to fit in the low 32 bits
139 if (packed && !fShapeX[fAttrAxis].isParam && fShapeX[fAttrAxis].dim > 0xFFFFFFFFULL)
140 packed = false;
141
142 std::string pairType = "std::pair<" + fType + ",int64_t>";
143 if (packed) {
144 out << SP << "std::vector<uint64_t> elements(" << n_elements << ");\n";
145 if (fShapeX[fAttrAxis].isParam) {
146 out << SP << "if (static_cast<unsigned long long>(" << n_elements << ") > 0xFFFFFFFFULL)\n";
147 out << SP << SP << "throw std::runtime_error(\"TMVA SOFIE TopK - reduced axis is longer "
148 << "than the 2^32 limit of the packed index\");\n";
149 }
150 } else {
151 out << SP << "std::vector<" << pairType << "> elements(" << n_elements << ");\n";
152 // taking the pairs by const reference avoids copying them on every comparison
153 out << SP << "auto " << OpName << "_cmp = [](const " << pairType << " &a, const " << pairType << " &b) {\n";
154 out << SP << SP << "return (a.first != b.first) ? (a.first " << (fAttrLargest ? ">" : "<")
155 << " b.first) : a.second < b.second;\n";
156 out << SP << "};\n";
157 }
158 // loop on elements before
159 if (n_before != "1") {
160 out << SP << "for (size_t i = 0; i < " << n_before << "; i++) {\n";
161 out << SP << SP << "size_t xoffset = i*" << strideX[axis-1] << ";\n";
162 out << SP << SP << "size_t yoffset = i*" << strideY[axis-1] << ";\n";
163 out << SP;
164 } else {
165 out << SP << "size_t xoffset = 0;\n";
166 out << SP << "size_t yoffset = 0;\n";
167 }
168 if (n_after != "1")
169 out << SP << "for (size_t j = 0; j < " << n_after << "; j++) {\n";
170 else
171 out << SP << "const size_t j = 0;\n";
172
173 // copy the elements to be sorted into the working buffer
174 out << SP << SP << "for (size_t l = 0; l < " << n_elements << "; l++) {\n";
175 if (packed) {
176 out << SP << SP << SP << "uint32_t b_ = 0;\n";
177 out << SP << SP << SP << "std::memcpy(&b_, &tensor_" << fNX << "[xoffset + " << strideX[axis]
178 << "*l + j], sizeof(b_));\n";
179 out << SP << SP << SP << "b_ ^= (b_ & 0x80000000u) ? 0xFFFFFFFFu : 0x80000000u;\n";
180 if (fAttrLargest)
181 out << SP << SP << SP << "b_ = ~b_;\n"; // reverse the value order, keep index ascending
182 out << SP << SP << SP << "elements[l] = (static_cast<uint64_t>(b_) << 32) | static_cast<uint32_t>(l);\n";
183 } else {
184 out << SP << SP << SP << "elements[l] = std::make_pair(tensor_" << fNX << "[xoffset + " << strideX[axis]
185 << "*l + j], l);\n";
186 }
187 out << SP << SP << "}\n";
188
189 // Move the K selected elements to the front in linear time, then order just those.
190 // std::partial_sort would be O(n log K) with heap operations over the whole range.
191 std::string cmp = packed ? "" : (", " + OpName + "_cmp");
192 out << SP << SP << "std::nth_element(elements.begin(), elements.begin() + (" << fK << "), elements.end()" << cmp
193 << ");\n";
194 // The ONNX spec leaves the order unspecified when sorted=0, but we sort anyway: it is
195 // only O(K log K) and it keeps the generated code reproducible across standard libraries.
196 out << SP << SP << "std::sort(elements.begin(), elements.begin() + (" << fK << ")" << cmp << ");\n";
197
198 // copy the selected elements in the output
199 out << SP << SP << "for (size_t l = 0; l < " << fK << "; l++) {\n";
200 if (packed) {
201 out << SP << SP << SP << "uint32_t b_ = static_cast<uint32_t>(elements[l] >> 32);\n";
202 if (fAttrLargest)
203 out << SP << SP << SP << "b_ = ~b_;\n";
204 out << SP << SP << SP << "b_ ^= (b_ & 0x80000000u) ? 0x80000000u : 0xFFFFFFFFu;\n";
205 out << SP << SP << SP << fType << " v_;\n";
206 out << SP << SP << SP << "std::memcpy(&v_, &b_, sizeof(v_));\n";
207 out << SP << SP << SP << "tensor_" << fNVal << "[yoffset + " << strideY[axis] << "*l + j] = v_;\n";
208 out << SP << SP << SP << "tensor_" << fNInd << "[yoffset + " << strideY[axis]
209 << "*l + j] = static_cast<int64_t>(static_cast<uint32_t>(elements[l]));\n";
210 } else {
211 out << SP << SP << SP << "tensor_" << fNVal << "[yoffset + " << strideY[axis]
212 << "*l + j] = elements[l].first;\n";
213 out << SP << SP << SP << "tensor_" << fNInd << "[yoffset + " << strideY[axis]
214 << "*l + j] = elements[l].second;\n";
215 }
216 out << SP << SP << "}\n";
217 if (n_after != "1") out << SP << SP << "}\n";
218 if (n_before != "1") out << SP << "}\n";
219 out << SP << "}\n"; // end operator scope
220 return out.str();
221 }
222};
223
224} // namespace SOFIE
225} // namespace Experimental
226} // namespace TMVA
227
228#endif // TMVA_SOFIE_ROPERATOR_TOPK
size_t size(const MatrixT &matrix)
retrieve the size of a square matrix
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t Float_t Float_t Float_t Int_t Int_t UInt_t UInt_t Rectangle_t Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char Pixmap_t Pixmap_t PictureAttributes_t attr const char char ret_data h unsigned char height h length
std::string Generate(std::string OpName) override
ROperator_TopK(int attr_axis, int attr_largest, int attr_sorted, std::string nameK, std::string nameX, std::string nameVal, std::string nameInd)
void Initialize(RModel &model) override
std::vector< std::string_view > fInputTensorNames
Definition ROperator.hxx:44
const std::string SP
space used to correctly indent the generated C++ code
Definition ROperator.hxx:40
std::vector< std::string_view > fOutputTensorNames
Definition ROperator.hxx:45
std::vector< size_t > ComputeStrideFromShape(const std::vector< size_t > &shape)
compute stride of a tensor given its shape (assume layout is row-major)
std::string ConvertDimShapeToString(const std::vector< Dim > &shape)
std::string ConvertTypeToString(ETensorType type)
std::string ConvertDimShapeToLength(const std::vector< Dim > &shape)
create variable transformations