ls1-MarDyn
ls1-MarDyn molecular dynamics code
TraversalTuner.h
1#ifndef TRAVERSALTUNER_H_
2#define TRAVERSALTUNER_H_
3
4#include <algorithm>
5#include <utility>
6#include <vector>
7
8#include <utils/Logger.h>
9#include <Simulation.h>
10#include "LinkedCellTraversals/CellPairTraversals.h"
11#include "LinkedCellTraversals/QuickschedTraversal.h"
12#include "LinkedCellTraversals/C08CellPairTraversal.h"
13#include "LinkedCellTraversals/C04CellPairTraversal.h"
14#include "LinkedCellTraversals/OriginalCellPairTraversal.h"
15#include "LinkedCellTraversals/HalfShellTraversal.h"
16#include "LinkedCellTraversals/MidpointTraversal.h"
17#include "LinkedCellTraversals/NeutralTerritoryTraversal.h"
18#include "LinkedCellTraversals/SlicedCellPairTraversal.h"
19
20using Log::global_log;
21
22template<class CellTemplate>
24 friend class LinkedCellsTest;
25
26public:
27 // Probably remove this once autotuning is implemented
28 enum traversalNames {
29 ORIGINAL = 0,
30 C08 = 1,
31 C04 = 2,
32 SLICED = 3,
33 HS = 4,
34 MP = 5,
35 C08ES = 6,
36 NT = 7,
37 // quicksched has to be the last traversal!
38 QSCHED = 8,
39 };
40
42
44
45 void findOptimalTraversal();
46
47 void readXML(XMLfileUnits &xmlconfig);
48
56 void rebuild(std::vector<CellTemplate> &cells,
57 const std::array<unsigned long, 3> &dims, double cellLength[3], double cutoff);
58
59 void traverseCellPairs(CellProcessor &cellProcessor);
60
61 void traverseCellPairs(traversalNames name, CellProcessor &cellProcessor);
62
63 void traverseCellPairsOuter(CellProcessor &cellProcessor);
64
65 void traverseCellPairsInner(CellProcessor &cellProcessor, unsigned stage, unsigned stageCount);
66
67
68 bool isTraversalApplicable(traversalNames name, const std::array<unsigned long, 3> &dims) const; // new
69
70
71 traversalNames getSelectedTraversal() const {
72 return selectedTraversal;
73 }
74
75 CellPairTraversals<ParticleCell> *getCurrentOptimalTraversal() { return _optimalTraversal; }
76
77private:
78 std::vector<CellTemplate>* _cells;
79 std::array<unsigned long, 3> _dims;
80
81 traversalNames selectedTraversal;
82
83 std::vector<std::pair<CellPairTraversals<CellTemplate> *, CellPairTraversalData *> > _traversals;
84
85 CellPairTraversals<CellTemplate> *_optimalTraversal;
86
87 unsigned _cellsInCutoff = 1;
88};
89
90template<class CellTemplate>
91TraversalTuner<CellTemplate>::TraversalTuner() : _cells(nullptr), _dims(), _optimalTraversal(nullptr) {
92 // defaults:
93 selectedTraversal = {
94 mardyn_get_max_threads() > 1 ? C08 : SLICED
95 };
96 auto *c08Data = new C08CellPairTraversalData;
97 auto *c04Data = new C04CellPairTraversalData;
98 auto *origData = new OriginalCellPairTraversalData;
99 auto *slicedData = new SlicedCellPairTraversalData;
100 auto *hsData = new HalfShellTraversalData;
101 auto *mpData = new MidpointTraversalData;
102 auto *ntData = new NeutralTerritoryTraversalData;
103 auto *c08esData = new C08CellPairTraversalData;
104
105 _traversals = {
106 make_pair(nullptr, origData),
107 make_pair(nullptr, c08Data),
108 make_pair(nullptr, c04Data),
109 make_pair(nullptr, slicedData),
110 make_pair(nullptr, hsData),
111 make_pair(nullptr, mpData),
112 make_pair(nullptr, ntData),
113 make_pair(nullptr, c08esData)
114 };
115#ifdef QUICKSCHED
117 quiData->taskBlockSize = {{2, 2, 2}};
118 if (is_base_of<ParticleCellBase, CellTemplate>::value) {
119 _traversals.push_back(make_pair(nullptr, quiData));
120 }
121#endif
122}
123
124template<class CellTemplate>
126 for (auto t : _traversals) {
127 if (t.first != nullptr)
128 delete (t.first);
129 if (t.second != nullptr)
130 delete (t.second);
131 }
132}
133
134template<class CellTemplate>
136 // TODO implement autotuning here! At the moment the traversal is chosen via readXML!
137
138 _optimalTraversal = _traversals[selectedTraversal].first;
139
140 // log traversal
141 if (dynamic_cast<HalfShellTraversal<CellTemplate> *>(_optimalTraversal))
142 global_log->info() << "Using HalfShellTraversal." << endl;
143 else if (dynamic_cast<OriginalCellPairTraversal<CellTemplate> *>(_optimalTraversal))
144 global_log->info() << "Using OriginalCellPairTraversal." << endl;
145 else if (dynamic_cast<C08CellPairTraversal<CellTemplate> *>(_optimalTraversal))
146 global_log->info() << "Using C08CellPairTraversal without eighthShell." << endl;
147 else if (dynamic_cast<C08CellPairTraversal<CellTemplate, true> *>(_optimalTraversal))
148 global_log->info() << "Using C08CellPairTraversal with eighthShell." << endl;
149 else if (dynamic_cast<C04CellPairTraversal<CellTemplate> *>(_optimalTraversal))
150 global_log->info() << "Using C04CellPairTraversal." << endl;
151 else if (dynamic_cast<MidpointTraversal<CellTemplate> *>(_optimalTraversal))
152 global_log->info() << "Using MidpointTraversal." << endl;
153 else if (dynamic_cast<NeutralTerritoryTraversal<CellTemplate> *>(_optimalTraversal))
154 global_log->info() << "Using NeutralTerritoryTraversal." << endl;
155 else if (dynamic_cast<QuickschedTraversal<CellTemplate> *>(_optimalTraversal)) {
156 global_log->info() << "Using QuickschedTraversal." << endl;
157#ifndef QUICKSCHED
158 global_log->error() << "MarDyn was compiled without Quicksched Support. Aborting!" << endl;
160#endif
161 } else if (dynamic_cast<SlicedCellPairTraversal<CellTemplate> *>(_optimalTraversal))
162 global_log->info() << "Using SlicedCellPairTraversal." << endl;
163 else
164 global_log->warning() << "Using unknown traversal." << endl;
165
166 if (_cellsInCutoff > _optimalTraversal->maxCellsInCutoff()) {
167 global_log->error() << "Traversal supports up to " << _optimalTraversal->maxCellsInCutoff()
168 << " cells in cutoff, but value is chosen as " << _cellsInCutoff << std::endl;
170 }
171}
172
173template<class CellTemplate>
175 string oldPath(xmlconfig.getcurrentnodepath());
176 // read traversal type default values
177 string traversalType;
178
179 xmlconfig.getNodeValue("traversalSelector", traversalType);
180 transform(traversalType.begin(), traversalType.end(), traversalType.begin(), ::tolower);
181
182 if (traversalType.find("c08es") != string::npos)
183 selectedTraversal = C08ES;
184 else if (traversalType.find("c08") != string::npos)
185 selectedTraversal = C08;
186 else if (traversalType.find("c04") != string::npos)
187 selectedTraversal = C04;
188 else if (traversalType.find("qui") != string::npos)
189 selectedTraversal = QSCHED;
190 else if (traversalType.find("slice") != string::npos)
191 selectedTraversal = SLICED;
192 else if (traversalType.find("ori") != string::npos)
193 selectedTraversal = ORIGINAL;
194 else if (traversalType.find("hs") != string::npos)
195 selectedTraversal = HS;
196 else if (traversalType.find("mp") != string::npos)
197 selectedTraversal = MP;
198 else if (traversalType.find("nt") != string::npos) {
199 selectedTraversal = NT;
200 } else {
201 // selector already set in constructor, just print a warning here
202 if (mardyn_get_max_threads() > 1) {
203 global_log->warning() << "No traversal type selected. Defaulting to c08 traversal." << endl;
204 } else {
205 global_log->warning() << "No traversal type selected. Defaulting to sliced traversal." << endl;
206 }
207 }
208
209 _cellsInCutoff = xmlconfig.getNodeValue_int("cellsInCutoffRadius", 1); // This is currently only used for an assert
210
211 // workaround for stupid iterator:
212 // since
213 // xmlconfig.changecurrentnode(traversalIterator);
214 // does not work resolve paths to traversals manually
215 // use iterator only to resolve number of traversals (==iterations)
216 string basePath(xmlconfig.getcurrentnodepath());
217
218 int i = 1;
219 XMLfile::Query qry = xmlconfig.query("traversalData");
220 for (XMLfile::Query::const_iterator traversalIterator = qry.begin(); traversalIterator; ++traversalIterator) {
221 string path(basePath + "/traversalData[" + to_string(i) + "]");
222 xmlconfig.changecurrentnode(path);
223
224 traversalType = xmlconfig.getNodeValue_string("@type", "NOTHING FOUND");
225 transform(traversalType.begin(), traversalType.end(), traversalType.begin(), ::tolower);
226 if (traversalType == "c08") {
227 // nothing to do
228 } else if (traversalType.find("qui") != string::npos) {
229#ifdef QUICKSCHED
230 if (not is_base_of<ParticleCellBase, CellTemplate>::value) {
231 global_log->warning() << "Attempting to use Quicksched with cell type that does not store task data!"
232 << endl;
233 }
234 for (auto p : _traversals) {
235 if (struct QuickschedTraversalData *quiData = dynamic_cast<QuickschedTraversalData *>(p.second)) {
236 // read task block size
237 string tag = "taskBlockSize/l";
238 char dimension = 'x';
239
240 for (int j = 0; j < 3; ++j) {
241 tag += (dimension + j);
242 xmlconfig.getNodeValue(tag, quiData->taskBlockSize[j]);
243 if (quiData->taskBlockSize[j] < 2) {
244 global_log->error() << "Task block size in "
245 << (char) (dimension + j)
246 << " direction is <2 and thereby invalid! ("
247 << quiData->taskBlockSize[j] << ")"
248 << endl;
250 }
251 }
252 break;
253 }
254 }
255#else
256 global_log->warning() << "Found quicksched traversal data in config "
257 << "but mardyn was compiled without quicksched support! "
258 << "(make ENABLE_QUICKSCHED=1)" << endl;
259#endif
260 } else {
261 global_log->warning() << "Unknown traversal type: " << traversalType << endl;
262 }
263 ++i;
264 }
265 xmlconfig.changecurrentnode(oldPath);
266}
267
268template<class CellTemplate>
269void TraversalTuner<CellTemplate>::rebuild(std::vector<CellTemplate> &cells, const std::array<unsigned long, 3> &dims,
270 double cellLength[3], double cutoff) {
271 _cells = &cells; // new - what for?
272 _dims = dims; // new - what for?
273
274 for (size_t i = 0ul; i < _traversals.size(); ++i) {
275 auto& [traversalPointerReference, traversalData] = _traversals[i];
276 // decide whether to initialize or rebuild
277 if (traversalPointerReference == nullptr) {
278 switch (i) {
279 case traversalNames ::ORIGINAL:
280 traversalPointerReference = new OriginalCellPairTraversal<CellTemplate>(cells, dims);
281 break;
282 case traversalNames::C08:
283 traversalPointerReference = new C08CellPairTraversal<CellTemplate>(cells, dims);
284 break;
285 case traversalNames::C04:
286 traversalPointerReference = new C04CellPairTraversal<CellTemplate>(cells, dims);
287 break;
288 case traversalNames::SLICED:
289 traversalPointerReference = new SlicedCellPairTraversal<CellTemplate>(cells, dims);
290 break;
291 case traversalNames::HS:
292 traversalPointerReference = new HalfShellTraversal<CellTemplate>(cells, dims);
293 break;
294 case traversalNames::MP:
295 traversalPointerReference = new MidpointTraversal<CellTemplate>(cells, dims);
296 break;
297 case traversalNames::NT:
298 traversalPointerReference =
299 new NeutralTerritoryTraversal<CellTemplate>(cells, dims, cellLength, cutoff);
300 break;
301 case traversalNames::C08ES:
302 traversalPointerReference = new C08CellPairTraversal<CellTemplate, true>(cells, dims);
303 break;
304 case traversalNames::QSCHED: {
305 mardyn_assert((is_base_of<ParticleCellBase, CellTemplate>::value));
306 auto *quiData = dynamic_cast<QuickschedTraversalData *>(traversalData);
307 traversalPointerReference = new QuickschedTraversal<CellTemplate>(cells, dims, quiData->taskBlockSize);
308 } break;
309 default:
310 global_log->error() << "Unknown traversal data found in TraversalTuner._traversals!" << endl;
312 }
313 }
314 traversalPointerReference->rebuild(cells, dims, cellLength, cutoff, traversalData);
315 }
316 _optimalTraversal = nullptr;
317}
318
319template<class CellTemplate>
321 if (not _optimalTraversal) {
322 findOptimalTraversal();
323 }
324 _optimalTraversal->traverseCellPairs(cellProcessor);
325}
326
327template<class CellTemplate>
329 CellProcessor& cellProcessor) {
330 if (name == getSelectedTraversal()) {
331 traverseCellPairs(cellProcessor);
332 } else {
333 SlicedCellPairTraversal<CellTemplate> slicedTraversal(*_cells, _dims);
334 switch(name) {
335 case SLICED:
336 slicedTraversal.traverseCellPairs(cellProcessor);
337 break;
338 default:
339 Log::global_log->error()<< "Calling traverseCellPairs(traversalName, CellProcessor&) for something else than the Sliced Traversal is disabled for now. Aborting." << std::endl;
340 mardyn_exit(1);
341 break;
342 }
343 }
344}
345
346template<class CellTemplate>
348 if (not _optimalTraversal) {
349 findOptimalTraversal();
350 }
351 _optimalTraversal->traverseCellPairsOuter(cellProcessor);
352}
353
354template<class CellTemplate>
356 unsigned stageCount) {
357 if (not _optimalTraversal) {
358 findOptimalTraversal();
359 }
360 _optimalTraversal->traverseCellPairsInner(cellProcessor, stage, stageCount);
361}
362
363template<class CellTemplate>
365 traversalNames name, const std::array<unsigned long, 3> &dims) const {
366 bool ret = true;
367 switch(name) {
368 case SLICED:
370 break;
371 case QSCHED:
372#ifdef QUICKSCHED
373 ret = true;
374#else
375 ret = false;
376#endif
377 break;
378 case C08:
379 ret = true;
380 break;
381 case C04:
382 ret = true;
383 break;
384 case ORIGINAL:
385 ret = true;
386 break;
387 default:
388 global_log->warning() << "unknown traversal given in TraversalTuner::isTraversalApplicable, assuming that is applicable" << std::endl;
389 }
390 return ret;
391}
392
393#endif //TRAVERSALTUNER_H_
Definition: C04CellPairTraversal.h:20
Definition: C08CellPairTraversal.h:20
Definition: CellPairTraversals.h:21
Definition: CellProcessor.h:29
Definition: HalfShellTraversal.h:17
Definition: MidpointTraversal.h:18
Definition: NeutralTerritoryTraversal.h:25
Definition: OriginalCellPairTraversal.h:21
Definition: QuickschedTraversal.h:27
static void exit(int exitcode)
Terminate simulation with given exit code.
Definition: Simulation.cpp:155
Definition: SlicedCellPairTraversal.h:19
Definition: TraversalTuner.h:23
void rebuild(std::vector< CellTemplate > &cells, const std::array< unsigned long, 3 > &dims, double cellLength[3], double cutoff)
Definition: TraversalTuner.h:269
XML file with unit attributes abstraction.
Definition: xmlfileUnits.h:25
Definition: xmlfile.h:240
Definition: xmlfile.h:232
const_iterator begin() const
get starting iterator return an iterator to the first node
Definition: xmlfile.h:344
Query query(const std::string &querystr) const
perform a query return a query to a given query expression
long changecurrentnode(const std::string &nodepath=std::string("/"))
set current node set a node, relative queries start with
int getNodeValue_int(const std::string &nodepath, int defaultvalue=0) const
get node value as int get the node content and convert it to an integer
Definition: xmlfile.h:514
std::string getNodeValue_string(const std::string &nodepath, const std::string defaultvalue=std::string()) const
get node value as string get the node content
Definition: xmlfile.h:499
unsigned long getNodeValue(const std::string &nodepath, T &value) const
get node value get the node content and convert it to a given type
Definition: xmlfile.h:484
std::string getcurrentnodepath() const
get current node path
Definition: xmlfile.h:477
::xsd::cxx::tree::name< char, token > name
C++ type corresponding to the Name XML Schema built-in type.
Definition: vtk-punstructured.h:288
Definition: C04CellPairTraversal.h:16
Definition: C08CellPairTraversal.h:16
Definition: CellPairTraversals.h:16
Definition: HalfShellTraversal.h:13
Definition: MidpointTraversal.h:14
Definition: NeutralTerritoryTraversal.h:13
Definition: OriginalCellPairTraversal.h:17
Definition: QuickschedTraversal.h:22
Definition: SlicedCellPairTraversal.h:15