TTK
Loading...
Searching...
No Matches
MergeTreeTemporalReduction.h
Go to the documentation of this file.
1
15
26
27#pragma once
28
29// ttk common includes
30#include <Debug.h>
31
32#include <FTMTreeUtils.h>
33#include <MergeTreeBarycenter.h>
34#include <MergeTreeBase.h>
35#include <MergeTreeDistance.h>
36#include <PathMappingDistance.h>
37
38namespace ttk {
39
44 class MergeTreeTemporalReduction : virtual public Debug,
45 public MergeTreeBase {
46 protected:
47 double removalPercentage_ = 50.;
48 bool usePathMappings_ = false;
49 bool useL2Distance_ = false;
50 std::vector<std::vector<double>> fieldL2_;
52 std::vector<double> timeVariable_;
53
54 public:
56
57 void setRemovalPercentage(double rs) {
59 }
60
61 void setUseL2Distance(bool useL2) {
62 useL2Distance_ = useL2;
63 }
64
65 void setPathMappings(bool usePM) {
66 usePathMappings_ = usePM;
67 }
68
69 template <class dataType>
70 dataType computeL2Distance(std::vector<dataType> &img1,
71 std::vector<dataType> &img2,
72 bool emptyFieldDistance = false) {
73 size_t const noPoints = img1.size();
74
75 std::vector<dataType> secondField = img2;
76 if(emptyFieldDistance)
77 secondField = std::vector<dataType>(noPoints, 0);
78
79 dataType distance = 0;
80
81 for(size_t i = 0; i < noPoints; ++i)
82 distance += std::pow((img1[i] - secondField[i]), 2);
83
84 distance = std::sqrt(distance);
85
86 return distance;
87 }
88
89 template <class dataType>
90 std::vector<dataType> computeL2Barycenter(std::vector<dataType> &img1,
91 std::vector<dataType> &img2,
92 double alpha) {
93
94 size_t const noPoints = img1.size();
95
96 std::vector<dataType> barycenter(noPoints);
97 for(size_t i = 0; i < noPoints; ++i)
98 barycenter[i] = alpha * img1[i] * (1 - alpha) * img2[i];
99
100 return barycenter;
101 }
102
103 template <class dataType>
106 bool emptyTreeDistance = false) {
107 dataType distance;
108 if(usePathMappings_) {
109 PathMappingDistance mergeTreeDistance;
110 mergeTreeDistance.setAssignmentSolver(assignmentSolverID_);
111 mergeTreeDistance.setEpsilonTree1(epsilonTree1_);
112 mergeTreeDistance.setEpsilonTree2(epsilonTree2_);
114 mergeTreeDistance.setThreadNumber(this->threadNumber_);
115 mergeTreeDistance.setDistanceSquaredRoot(false);
116 mergeTreeDistance.setDebugLevel(2); // ToDo: Why??
117 mergeTreeDistance.setPreprocess(false);
118 mergeTreeDistance.setComputeMapping(false);
119
120 ftm::FTMTree_MT *mt1 = &(mTree1.tree);
121 ftm::FTMTree_MT *mt2 = &(mTree2.tree);
122 distance = mergeTreeDistance.computeDistance<dataType>(mt1, mt2);
123 } else {
124 MergeTreeDistance mergeTreeDistance;
125 mergeTreeDistance.setAssignmentSolver(assignmentSolverID_);
126 mergeTreeDistance.setEpsilonTree1(epsilonTree1_);
127 mergeTreeDistance.setEpsilonTree2(epsilonTree2_);
128 mergeTreeDistance.setEpsilon2Tree1(epsilon2Tree1_);
129 mergeTreeDistance.setEpsilon2Tree2(epsilon2Tree2_);
130 mergeTreeDistance.setEpsilon3Tree1(epsilon3Tree1_);
131 mergeTreeDistance.setEpsilon3Tree2(epsilon3Tree2_);
133 mergeTreeDistance.setParallelize(parallelize_);
136 mergeTreeDistance.setKeepSubtree(keepSubtree_);
137 mergeTreeDistance.setUseMinMaxPair(useMinMaxPair_);
138 mergeTreeDistance.setThreadNumber(this->threadNumber_);
139 mergeTreeDistance.setDistanceSquaredRoot(true); // squared root
140 mergeTreeDistance.setDebugLevel(2);
141 mergeTreeDistance.setPreprocess(false);
142 mergeTreeDistance.setPostprocess(false);
143 // mergeTreeDistance.setIsCalled(true);
144 mergeTreeDistance.setOnlyEmptyTreeDistance(emptyTreeDistance);
145
146 std::vector<std::tuple<ftm::idNode, ftm::idNode, double>> matching;
147 distance
148 = mergeTreeDistance.execute<dataType>(mTree1, mTree2, matching);
149 }
150
151 return distance;
152 }
153
154 template <class dataType>
157 double alpha) {
158 MergeTreeBarycenter mergeTreeBarycenter;
159 mergeTreeBarycenter.setAssignmentSolver(assignmentSolverID_);
160 mergeTreeBarycenter.setEpsilonTree1(epsilonTree1_);
161 mergeTreeBarycenter.setEpsilonTree2(epsilonTree2_);
162 mergeTreeBarycenter.setEpsilon2Tree1(epsilon2Tree1_);
163 mergeTreeBarycenter.setEpsilon2Tree2(epsilon2Tree2_);
164 mergeTreeBarycenter.setEpsilon3Tree1(epsilon3Tree1_);
165 mergeTreeBarycenter.setEpsilon3Tree2(epsilon3Tree2_);
166 mergeTreeBarycenter.setBranchDecomposition(branchDecomposition_);
167 mergeTreeBarycenter.setParallelize(parallelize_);
169 mergeTreeBarycenter.setKeepSubtree(keepSubtree_);
170 mergeTreeBarycenter.setUseMinMaxPair(useMinMaxPair_);
171 mergeTreeBarycenter.setThreadNumber(this->threadNumber_);
172 mergeTreeBarycenter.setAlpha(alpha);
173 mergeTreeBarycenter.setDebugLevel(2);
174 mergeTreeBarycenter.setPreprocess(false);
175 mergeTreeBarycenter.setPostprocess(false);
176 // mergeTreeBarycenter.setIsCalled(true);
177
178 if(usePathMappings_) {
179 mergeTreeBarycenter.setBaseModule(2);
180 mergeTreeBarycenter.setBranchDecomposition(false);
181 mergeTreeBarycenter.setNormalizedWasserstein(false);
182 mergeTreeBarycenter.setKeepSubtree(false);
183 mergeTreeBarycenter.setUseMinMaxPair(false);
184 mergeTreeBarycenter.setAddNodes(false);
185 mergeTreeBarycenter.setPostprocess(false);
186 } else {
187 mergeTreeBarycenter.setBranchDecomposition(true);
189 // mergeTreeBarycenter.setNormalizedWassersteinReg(normalizedWassersteinReg_);
190 // mergeTreeBarycenter.setRescaledWasserstein(rescaledWasserstein_);
191 }
192
193 std::vector<ftm::MergeTree<dataType>> intermediateTrees;
194 intermediateTrees.push_back(mTree1);
195 intermediateTrees.push_back(mTree2);
196 std::vector<std::vector<std::tuple<ftm::idNode, ftm::idNode, double>>>
197 outputMatchingBarycenter(2);
198 ftm::MergeTree<dataType> barycenter;
199 mergeTreeBarycenter.execute<dataType>(
200 intermediateTrees, outputMatchingBarycenter, barycenter);
201 return barycenter;
202 }
203
204 double computeAlpha(int index1, int middleIndex, int index2) {
205 index1 = timeVariable_[index1];
206 middleIndex = timeVariable_[middleIndex];
207 index2 = timeVariable_[index2];
208 return 1 - ((double)middleIndex - index1) / (index2 - index1);
209 }
210
211 template <class dataType>
212 void
214 std::vector<int> &removed,
215 std::vector<ftm::MergeTree<dataType>> &barycenters,
216 std::vector<std::vector<dataType>> &barycentersL2) {
217 std::vector<bool> treeRemoved(mTrees.size(), false);
218
219 int toRemoved = mTrees.size() * removalPercentage_ / 100.;
220 toRemoved = std::min(toRemoved, (int)(mTrees.size() - 2));
221
222 std::vector<std::vector<dataType>> images(fieldL2_.size());
223 for(size_t i = 0; i < fieldL2_.size(); ++i)
224 for(size_t j = 0; j < fieldL2_[i].size(); ++j)
225 images[i].push_back(static_cast<dataType>(fieldL2_[i][j]));
226
227 for(int iter = 0; iter < toRemoved; ++iter) {
228 dataType bestCost = std::numeric_limits<dataType>::max();
229 int bestMiddleIndex = -1;
230 ftm::MergeTree<dataType> bestBarycenter;
231 std::vector<std::tuple<ftm::MergeTree<dataType>, int>>
232 bestBarycentersOnPath;
233 std::vector<dataType> bestBarycenterL2;
234 std::vector<std::tuple<std::vector<dataType>, int>>
235 bestBarycentersL2OnPath;
236
237 // Compute barycenter for each pair of trees
238 printMsg("Compute barycenter for each pair of trees",
240 unsigned int index1 = 0, index2 = 0;
241 while(index2 != mTrees.size() - 1) {
242
243 // Get index in the middle
244 int middleIndex = index1 + 1;
245 while(treeRemoved[middleIndex])
246 ++middleIndex;
247
248 // Get second index
249 index2 = middleIndex + 1;
250 while(treeRemoved[index2])
251 ++index2;
252
253 // Compute barycenter
254 printMsg("Compute barycenter", debug::Priority::VERBOSE);
255 double const alpha = computeAlpha(index1, middleIndex, index2);
256 ftm::MergeTree<dataType> barycenter;
257 std::vector<dataType> barycenterL2;
258 if(not useL2Distance_)
259 barycenter = computeBarycenter<dataType>(
260 mTrees[index1], mTrees[index2], alpha);
261 else
262 barycenterL2 = computeL2Barycenter<dataType>(
263 images[index1], images[index2], alpha);
264
265 // - Compute cost
266 // Compute distance with middleIndex
267 printMsg(
268 "Compute distance with middleIndex", debug::Priority::VERBOSE);
269 dataType cost;
270 if(not useL2Distance_)
271 cost = computeDistance<dataType>(barycenter, mTrees[middleIndex]);
272 else
273 cost
274 = computeL2Distance<dataType>(barycenterL2, images[middleIndex]);
275
276 // Compute distances of previously removed trees on the path
277 printMsg("Compute distances of previously removed trees",
279 std::vector<std::tuple<ftm::MergeTree<dataType>, int>>
280 barycentersOnPath;
281 std::vector<std::tuple<std::vector<dataType>, int>>
282 barycentersL2OnPath;
283 for(unsigned int i = 0; i < 2; ++i) {
284 int const toReach = (i == 0 ? index1 : index2);
285 int const offset = (i == 0 ? -1 : 1);
286 int tIndex = middleIndex + offset;
287 while(tIndex != toReach) {
288
289 // Compute barycenter
290 double const alphaT = computeAlpha(index1, tIndex, index2);
291 ftm::MergeTree<dataType> barycenterP;
292 std::vector<dataType> barycenterPL2;
293 if(not useL2Distance_)
294 barycenterP = computeBarycenter<dataType>(
295 mTrees[index1], mTrees[index2], alphaT);
296 else
297 barycenterPL2 = computeL2Barycenter<dataType>(
298 images[index1], images[index2], alphaT);
299
300 // Compute distance
301 dataType costP;
302 if(not useL2Distance_)
303 costP = computeDistance<dataType>(barycenterP, mTrees[tIndex]);
304 else
305 costP
306 = computeL2Distance<dataType>(barycenterPL2, images[tIndex]);
307
308 // Save results
309 if(not useL2Distance_)
310 barycentersOnPath.push_back(
311 std::make_tuple(barycenterP, tIndex));
312 else
313 barycentersL2OnPath.push_back(
314 std::make_tuple(barycenterPL2, tIndex));
315 cost += costP;
316 tIndex += offset;
317 }
318 }
319
320 if(cost < bestCost) {
321 bestCost = cost;
322 bestMiddleIndex = middleIndex;
323 if(not useL2Distance_) {
324 bestBarycenter = barycenter;
325 bestBarycentersOnPath = barycentersOnPath;
326 } else {
327 bestBarycenterL2 = barycenterL2;
328 bestBarycentersL2OnPath = barycentersL2OnPath;
329 }
330 }
331
332 // Go to the next index
333 index1 = middleIndex;
334 }
335
336 // Removed the tree with the lowest cost
337 printMsg(
338 "Removed the tree with the lowest cost", debug::Priority::VERBOSE);
339 removed.push_back(bestMiddleIndex);
340 treeRemoved[bestMiddleIndex] = true;
341 if(not useL2Distance_) {
342 barycenters[bestMiddleIndex] = bestBarycenter;
343 for(auto &tup : bestBarycentersOnPath)
344 barycenters[std::get<1>(tup)] = std::get<0>(tup);
345 } else {
346 barycentersL2[bestMiddleIndex] = bestBarycenterL2;
347 for(auto &tup : bestBarycentersL2OnPath)
348 barycentersL2[std::get<1>(tup)] = std::get<0>(tup);
349 }
350 }
351 }
352
353 template <class dataType>
354 std::vector<int> execute(std::vector<ftm::MergeTree<dataType>> &mTrees,
355 std::vector<double> &emptyTreeDistances,
356 std::vector<ftm::MergeTree<dataType>> &allMT) {
357 Timer t_tempSub;
358
359 // --- Preprocessing
360 if(not useL2Distance_) { //} && not usePathMappings_) {
361 treesNodeCorr_ = std::vector<std::vector<int>>(mTrees.size());
362 for(unsigned int i = 0; i < mTrees.size(); ++i) {
366 true, usePathMappings_);
367 }
369 }
370 // if(usePathMappings_){
371 // treesNodeCorr_ = std::vector<std::vector<int>>(mTrees.size());
372 // printMsg("uses path mapping distance");
373 // for(unsigned int i = 0; i < mTrees.size(); ++i) {
374 // ftm::FTMTree_MT *tree = &(mTrees[i].tree);
375 // preprocessTree<dataType>(tree, true);
376
377 // // - Delete null persistence pairs and persistence thresholding
378 // persistenceThresholding<dataType>(tree, persistenceThreshold_);
379
380 // // - Merge saddle points according epsilon
381 // if(not isPersistenceDiagram_) {
382 // if(epsilonTree2_ != 0){
383 // std::vector<std::vector<ftm::idNode>> treeNodeMerged(
384 // tree->getNumberOfNodes() ); mergeSaddle<dataType>(tree,
385 // epsilonTree2_, treeNodeMerged); for(unsigned int j=0;
386 // j<treeNodeMerged.size(); j++){
387 // for(auto k : treeNodeMerged[j]){
388 // auto nodeToDelete = tree->getNode(j)->getOrigin();
389 // tree->getNode(k)->setOrigin(j);
390 // tree->getNode(nodeToDelete)->setOrigin(-1);
391 // }
392 // }
393 // ftm::cleanMergeTree<dataType>(mTrees[i], treesNodeCorr_[i],
394 // true);
395 // }
396 // else{
397 // std::vector<ttk::SimplexId>
398 // nodeCorri(tree->getNumberOfNodes()); for(unsigned int j=0;
399 // j<nodeCorri.size(); j++) nodeCorri[j] = j; treesNodeCorr_[i] =
400 // nodeCorri;
401 // }
402 // }
403 // if(deleteMultiPersPairs_)
404 // deleteMultiPersPairs<dataType>(tree, false);
405 // }
406 // }
407
408 // --- Execute
409 std::vector<ftm::MergeTree<dataType>> barycenters(mTrees.size());
410 std::vector<std::vector<dataType>> barycentersL2(mTrees.size());
411 std::vector<int> removed;
412 if(not useCustomTimeVariable_) {
413 timeVariable_.clear();
414 for(size_t i = 0; i < mTrees.size(); ++i)
415 timeVariable_.push_back(i);
416 }
418 mTrees, removed, barycenters, barycentersL2);
419
420 // --- Concatenate all trees/L2Images
421 std::vector<std::vector<dataType>> images(fieldL2_.size());
422 for(size_t i = 0; i < fieldL2_.size(); ++i)
423 for(size_t j = 0; j < fieldL2_[i].size(); ++j)
424 images[i].push_back(static_cast<dataType>(fieldL2_[i][j]));
425
426 for(auto &mt : mTrees)
427 allMT.push_back(mt);
428 std::vector<bool> removedB(mTrees.size(), false);
429 for(auto r : removed)
430 removedB[r] = true;
431 for(unsigned int i = 0; i < barycenters.size(); ++i)
432 if(removedB[i]) {
433 if(not useL2Distance_)
434 allMT.push_back(barycenters[i]);
435 else
436 images.push_back(barycentersL2[i]);
437 }
438
439 // --- Compute empty tree distances
440 unsigned int const distMatSize
441 = (not useL2Distance_ ? allMT.size() : images.size());
442 for(unsigned int i = 0; i < distMatSize; ++i) {
443 dataType distance;
444 if(not useL2Distance_)
445 distance = computeDistance<dataType>(allMT[i], allMT[i], true);
446 else
447 distance = computeL2Distance<dataType>(images[i], images[i], true);
448 emptyTreeDistances.push_back(distance);
449 }
450
451 // --- Postprocessing
452 if(not useL2Distance_ && not usePathMappings_) {
453 for(unsigned int i = 0; i < allMT.size(); ++i)
454 postprocessingPipeline<dataType>(&(allMT[i].tree));
455 for(unsigned int i = 0; i < mTrees.size(); ++i)
456 postprocessingPipeline<dataType>(&(mTrees[i].tree));
457 }
458
459 // --- Print results
460 std::stringstream ss, ss2, ss3;
461 ss << "input size = " << mTrees.size();
462 printMsg(ss.str());
463 ss2 << "output size = "
464 << mTrees.size() - (distMatSize - mTrees.size());
465 printMsg(ss2.str());
466 ss3 << "removed : ";
467 for(unsigned int i = 0; i < removed.size(); ++i) {
468 auto r = removed[i];
469 ss3 << r;
470 if(i < removed.size() - 1)
471 ss3 << ", ";
472 }
473 printMsg(ss3.str());
474
475 sort(removed.begin(), removed.end());
476
477 printMsg("Encoding", 1, t_tempSub.getElapsedTime(), this->threadNumber_);
478
479 return removed;
480 }
481
482 }; // MergeTreeTemporalReduction class
483
484} // namespace ttk
virtual int setThreadNumber(const int threadNumber)
Definition BaseClass.h:80
virtual int setDebugLevel(const int &debugLevel)
Definition Debug.cpp:147
void setAddNodes(bool addNodesT)
void setPreprocess(bool preproc)
void setPostprocess(bool postproc)
void execute(std::vector< ftm::MergeTree< dataType > > &trees, std::vector< double > &alphas, std::vector< std::vector< std::tuple< ftm::idNode, ftm::idNode, double > > > &finalMatchings, std::vector< std::vector< std::pair< std::pair< ftm::idNode, ftm::idNode >, std::pair< ftm::idNode, ftm::idNode > > > > &finalMatchings_path, ftm::MergeTree< dataType > &baryMergeTree, bool finalAsgnDoubleInput=false, bool finalAsgnFirstInput=true)
void setBranchDecomposition(bool useBD)
void setNormalizedWasserstein(bool normalizedWasserstein)
void setDistanceSquaredRoot(bool distanceSquaredRoot)
void setEpsilon3Tree1(double epsilon)
void setEpsilonTree1(double epsilon)
void setAssignmentSolver(int assignmentSolver)
void setEpsilon2Tree1(double epsilon)
void setEpsilonTree2(double epsilon)
void setPersistenceThreshold(double pt)
void preprocessingPipeline(ftm::MergeTree< dataType > &mTree, double epsilonTree, double epsilon2Tree, double epsilon3Tree, bool branchDecompositionT, bool useMinMaxPairT, bool cleanTreeT, double persistenceThreshold, std::vector< int > &nodeCorr, bool deleteInconsistentNodes=true, bool removeMergedSaddles=false)
void printTreesStats(std::vector< ftm::FTMTree_MT * > &trees)
void postprocessingPipeline(ftm::FTMTree_MT *tree)
std::vector< std::vector< int > > treesNodeCorr_
void setEpsilon2Tree2(double epsilon)
void setKeepSubtree(bool keepSubtree)
void setUseMinMaxPair(bool useMinMaxPair)
void setEpsilon3Tree2(double epsilon)
void setParallelize(bool para)
void setOnlyEmptyTreeDistance(double only)
void setPreprocess(bool preproc)
void setPostprocess(bool postproc)
dataType execute(ftm::MergeTree< dataType > &mTree1, ftm::MergeTree< dataType > &mTree2, std::vector< std::tuple< ftm::idNode, ftm::idNode, double > > &outputMatching)
dataType computeDistance(ftm::MergeTree< dataType > &mTree1, ftm::MergeTree< dataType > &mTree2, bool emptyTreeDistance=false)
double computeAlpha(int index1, int middleIndex, int index2)
std::vector< int > execute(std::vector< ftm::MergeTree< dataType > > &mTrees, std::vector< double > &emptyTreeDistances, std::vector< ftm::MergeTree< dataType > > &allMT)
std::vector< std::vector< double > > fieldL2_
ftm::MergeTree< dataType > computeBarycenter(ftm::MergeTree< dataType > &mTree1, ftm::MergeTree< dataType > &mTree2, double alpha)
void temporalSubsampling(std::vector< ftm::MergeTree< dataType > > &mTrees, std::vector< int > &removed, std::vector< ftm::MergeTree< dataType > > &barycenters, std::vector< std::vector< dataType > > &barycentersL2)
std::vector< dataType > computeL2Barycenter(std::vector< dataType > &img1, std::vector< dataType > &img2, double alpha)
dataType computeL2Distance(std::vector< dataType > &img1, std::vector< dataType > &img2, bool emptyFieldDistance=false)
void setAssignmentSolver(int assignmentSolver)
dataType computeDistance(ftm::FTMTree_MT *tree1, ftm::FTMTree_MT *tree2, std::vector< std::pair< std::pair< ftm::idNode, ftm::idNode >, std::pair< ftm::idNode, ftm::idNode > > > *outputMatching)
double getElapsedTime()
Definition Timer.h:15
TTK base package defining the standard types.
ftm::FTMTree_MT tree
Definition FTMTree_MT.h:906
printMsg(debug::output::BOLD+" | | | | | . \\ | | (__| | / __/| |_| / __/| (_) |"+debug::output::ENDCOLOR, debug::Priority::PERFORMANCE, debug::LineMode::NEW, stream)