TTK
Loading...
Searching...
No Matches
MergeTreeDistanceMatrix.h
Go to the documentation of this file.
1
47
48#pragma once
49
50// ttk common includes
51#include <Debug.h>
52
54#include <FTMTree.h>
55#include <FTMTreeUtils.h>
56#include <MergeTreeBase.h>
57#include <MergeTreeDistance.h>
58#include <PathMappingDistance.h>
59
60namespace ttk {
61
66 class MergeTreeDistanceMatrix : virtual public Debug,
67 virtual public MergeTreeBase {
68 protected:
69 int baseModule_ = 0;
71 int pathMetric_ = 0;
72
73 public:
76 "MergeTreeDistanceMatrix"); // inherited from Debug: prefix will be
77 // printed at the
78 // beginning of every msg
79 }
80 ~MergeTreeDistanceMatrix() override = default;
81
82 void setBaseModule(int m) {
83 baseModule_ = m;
84 }
85
86 void setBranchMetric(int m) {
87 branchMetric_ = m;
88 }
89
90 void setPathMetric(int m) {
91 pathMetric_ = m;
92 }
93
97 template <class dataType>
98 void execute(std::vector<ftm::MergeTree<dataType>> &trees,
99 std::vector<ftm::MergeTree<dataType>> &trees2,
100 std::vector<std::vector<double>> &distanceMatrix) {
101 treesNodeCorr_.resize(trees.size());
102 for(unsigned int i = 0; i < trees.size(); ++i) {
105 baseModule_ == 0 ? branchDecomposition_ : false, useMinMaxPair_, true,
106 treesNodeCorr_[i], true, baseModule_ == 2);
107 }
108 executePara<dataType>(trees, distanceMatrix);
109 if(trees2.size() != 0) {
110 std::vector<std::vector<int>> trees2NodeCorr(trees2.size());
111 for(unsigned int i = 0; i < trees.size(); ++i) {
115 true, treesNodeCorr_[i], true, baseModule_ == 2);
116 }
117 useDoubleInput_ = true;
118 std::vector<std::vector<double>> distanceMatrix2(
119 trees2.size(), std::vector<double>(trees2.size()));
120 executePara<dataType>(trees2, distanceMatrix2, false);
121 mixDistancesMatrix(distanceMatrix, distanceMatrix2);
122 }
123 }
124
125 template <class dataType>
126 void executePara(std::vector<ftm::MergeTree<dataType>> &trees,
127 std::vector<std::vector<double>> &distanceMatrix,
128 bool isFirstInput = true) {
129#ifdef TTK_ENABLE_OPENMP
130#pragma omp parallel num_threads(this->threadNumber_)
131 {
132#pragma omp single nowait
133#endif
134 executeParaImpl<dataType>(trees, distanceMatrix, isFirstInput);
135#ifdef TTK_ENABLE_OPENMP
136#pragma omp taskwait
137 } // pragma omp parallel
138#endif
139 }
140
141 template <class dataType>
143 std::vector<std::vector<double>> &distanceMatrix,
144 bool isFirstInput = true) {
145 for(unsigned int i = 0; i < distanceMatrix.size(); ++i) {
146#ifdef TTK_ENABLE_OPENMP
147#pragma omp task firstprivate(i) UNTIED() shared(distanceMatrix, trees)
148 {
149#endif
150 if(i % std::max(int(distanceMatrix.size() / 10), 1) == 0) {
151 std::stringstream stream;
152 stream << i << " / " << distanceMatrix.size();
153 printMsg(stream.str());
154 }
155 distanceMatrix[i][i] = 0.0;
156 for(unsigned int j = i + 1; j < distanceMatrix[0].size(); ++j) {
157 // Execute
158 if(baseModule_ == 0) {
159 MergeTreeDistance mergeTreeDistance;
160 mergeTreeDistance.setAssignmentSolver(assignmentSolverID_);
161 mergeTreeDistance.setEpsilonTree1(epsilonTree1_);
162 mergeTreeDistance.setEpsilonTree2(epsilonTree2_);
163 mergeTreeDistance.setEpsilon2Tree1(epsilon2Tree1_);
164 mergeTreeDistance.setEpsilon2Tree2(epsilon2Tree2_);
165 mergeTreeDistance.setEpsilon3Tree1(epsilon3Tree1_);
166 mergeTreeDistance.setEpsilon3Tree2(epsilon3Tree2_);
168 mergeTreeDistance.setParallelize(parallelize_);
170 mergeTreeDistance.setDebugLevel(std::min(debugLevel_, 2));
171 mergeTreeDistance.setThreadNumber(this->threadNumber_);
172 mergeTreeDistance.setNormalizedWasserstein(
174 mergeTreeDistance.setKeepSubtree(keepSubtree_);
176 mergeTreeDistance.setUseMinMaxPair(useMinMaxPair_);
177 mergeTreeDistance.setPreprocess(false);
178 // mergeTreeDistance.setSaveTree(true);
179 mergeTreeDistance.setSaveTree(false);
180 mergeTreeDistance.setCleanTree(true);
181 mergeTreeDistance.setIsCalled(true);
182 mergeTreeDistance.setPostprocess(false);
184 if(useDoubleInput_) {
185 double const weight
186 = mixDistancesMinMaxPairWeight(isFirstInput);
187 mergeTreeDistance.setMinMaxPairWeight(weight);
188 mergeTreeDistance.setDistanceSquaredRoot(true);
189 }
190 std::vector<std::tuple<ftm::idNode, ftm::idNode>> outputMatching;
191 distanceMatrix[i][j] = mergeTreeDistance.execute<dataType>(
192 trees[i], trees[j], outputMatching);
193 } else if(baseModule_ == 1) {
194 BranchMappingDistance branchDist;
195 branchDist.setBaseMetric(branchMetric_);
198 branchDist.setEpsilonTree1(epsilonTree1_);
199 branchDist.setEpsilonTree2(epsilonTree2_);
205 branchDist.setPreprocess(false);
206 // branchDist.setSaveTree(true);
207 branchDist.setSaveTree(false);
208 dataType dist = branchDist.execute<dataType>(trees[i], trees[j]);
209 distanceMatrix[i][j] = static_cast<double>(dist);
210 } else if(baseModule_ == 2) {
211 PathMappingDistance pathDist;
212 pathDist.setBaseMetric(pathMetric_);
215 pathDist.setComputeMapping(true);
223 pathDist.setPreprocess(false);
224 // pathDist.setSaveTree(true);
225 pathDist.setSaveTree(false);
226 dataType dist = pathDist.execute<dataType>(trees[i], trees[j]);
227 distanceMatrix[i][j] = static_cast<double>(dist);
228 }
229 // distance matrix is symmetric
230 distanceMatrix[j][i] = distanceMatrix[i][j];
231 } // end for j
232#ifdef TTK_ENABLE_OPENMP
233 } // end task
234#endif
235 } // end for i
236 }
237
238 }; // MergeTreeDistanceMatrix class
239
240} // namespace ttk
virtual int setThreadNumber(const int threadNumber)
Definition BaseClass.h:80
dataType execute(ftm::MergeTree< dataType > &mTree1, ftm::MergeTree< dataType > &mTree2, std::vector< std::tuple< ftm::idNode, ftm::idNode, double > > *outputMatching=nullptr)
void setAssignmentSolver(int assignmentSolver)
int debugLevel_
Definition Debug.h:379
void setDebugMsgPrefix(const std::string &prefix)
Definition Debug.h:364
virtual int setDebugLevel(const int &debugLevel)
Definition Debug.cpp:147
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)
std::vector< std::vector< int > > treesNodeCorr_
void setCleanTree(bool clean)
void setEpsilon2Tree2(double epsilon)
void setKeepSubtree(bool keepSubtree)
void setUseMinMaxPair(bool useMinMaxPair)
void setEpsilon3Tree2(double epsilon)
double mixDistancesMinMaxPairWeight(bool isFirstInput)
void setParallelize(bool para)
void setIsPersistenceDiagram(bool isPD)
void mixDistancesMatrix(std::vector< std::vector< dataType > > &distanceMatrix, std::vector< std::vector< dataType > > &distanceMatrix2)
void executePara(std::vector< ftm::MergeTree< dataType > > &trees, std::vector< std::vector< double > > &distanceMatrix, bool isFirstInput=true)
void executeParaImpl(std::vector< ftm::MergeTree< dataType > > &trees, std::vector< std::vector< double > > &distanceMatrix, bool isFirstInput=true)
void execute(std::vector< ftm::MergeTree< dataType > > &trees, std::vector< ftm::MergeTree< dataType > > &trees2, std::vector< std::vector< double > > &distanceMatrix)
~MergeTreeDistanceMatrix() override=default
void setPreprocess(bool preproc)
void setPostprocess(bool postproc)
void setSaveTree(bool save)
void setMinMaxPairWeight(double weight)
dataType execute(ftm::MergeTree< dataType > &mTree1, ftm::MergeTree< dataType > &mTree2, std::vector< std::tuple< ftm::idNode, ftm::idNode, double > > &outputMatching)
void setAssignmentSolver(int assignmentSolver)
dataType execute(ftm::MergeTree< dataType > &mTree1, ftm::MergeTree< dataType > &mTree2, std::vector< std::pair< std::pair< ftm::idNode, ftm::idNode >, std::pair< ftm::idNode, ftm::idNode > > > *outputMatching)
TTK base package defining the standard types.
printMsg(debug::output::BOLD+" | | | | | . \\ | | (__| | / __/| |_| / __/| (_) |"+debug::output::ENDCOLOR, debug::Priority::PERFORMANCE, debug::LineMode::NEW, stream)