Belle II Software development
WeightedFastHoughTree.h
1/**************************************************************************
2 * basf2 (Belle II Analysis Software Framework) *
3 * Author: The Belle II Collaboration *
4 * *
5 * See git log for contributors and copyright holders. *
6 * This file is licensed under LGPL-3.0, see LICENSE.md. *
7 **************************************************************************/
8#pragma once
9
10#include <tracking/trackFindingCDC/hough/trees/DynTree.h>
11#include <tracking/trackFindingCDC/hough/baseelements/WithWeightedItems.h>
12#include <tracking/trackFindingCDC/hough/baseelements/WithSharedMark.h>
13
14#include <vector>
15#include <memory>
16#include <cassert>
17#include <cfloat>
18#include <cmath>
19#include <algorithm>
20#include <type_traits>
21
22namespace Belle2 {
27 namespace TrackFindingCDC {
28
30 template<class T, class ADomain, class ADomainDivsion>
32 public DynTree< WithWeightedItems<ADomain, T>, ADomainDivsion> {
33 private:
36
37 public:
39 WeightedParititioningDynTree(ADomain topDomain, ADomainDivsion domainDivsion) :
40 Super(WithWeightedItems<ADomain, T>(std::move(topDomain)), std::move(domainDivsion))
41 {}
42 };
43
50 template<class T, class ADomain, class ADomainDivsion>
52 public WeightedParititioningDynTree<WithSharedMark<T>, ADomain, ADomainDivsion> {
53
54 private:
56 using Super = WeightedParititioningDynTree<WithSharedMark<T>, ADomain, ADomainDivsion>;
57
58 public:
61
63 using Node = typename Super::Node;
64
65 public:
67 template<class Ts>
68 void seed(const Ts& items)
69 {
70 this->fell();
71 Node& topNode = this->getTopNode();
72 for (auto&& item : items) {
73 m_marks.push_back(false);
74 bool& markOfItem = m_marks.back();
75 TrackingUtilities::Weight weight = DBL_MAX;
76 topNode.insert(WithSharedMark<T>(T(item), &markOfItem), weight);
77 }
78 }
79
81 template <class AItemInDomainMeasure>
82 std::vector<std::pair<ADomain, std::vector<T>>>
83 findHeavyLeavesDisjoint(AItemInDomainMeasure& weightItemInDomain,
84 int maxLevel,
85 double minWeight)
86 {
87 auto skipLowWeightNode = [minWeight](const Node * node) {
88 return not(node->getWeight() >= minWeight);
89 };
90 return findLeavesDisjoint(weightItemInDomain, maxLevel, skipLowWeightNode);
91 }
92
94 template <class AItemInDomainMeasure, class ASkipNodePredicate>
95 std::vector<std::pair<ADomain, std::vector<T>>>
96 findLeavesDisjoint(AItemInDomainMeasure& weightItemInDomain,
97 int maxLevel,
98 ASkipNodePredicate& skipNode)
99 {
100 std::vector<std::pair<ADomain, std::vector<T> > > found;
101 auto isLeaf = [&found, &skipNode, maxLevel](Node * node) {
102 // Skip the expansion and the filling of the children
103 if (skipNode(node)) {
104 return true;
105 }
106
107 // Node is a leaf at the maximum level
108 // Save its content
109 // Do not walk children
110 if (node->getLevel() >= maxLevel) {
111 const ADomain* domain = node;
112 found.emplace_back(*domain, std::vector<T>(node->begin(), node->end()));
113 for (WithSharedMark<T>& markableItem : *node) {
114 markableItem.mark();
115 }
116 return true;
117 }
118
119 // Else to node has enough weight and is not at the lowest level
120 // Signal that it is not a leaf
121 // Continue to create and fill children.
122 return false;
123 };
124 fillWalk(weightItemInDomain, isLeaf);
125 return found;
126 }
127
135 template <class AItemInDomainMeasure>
136 std::vector<std::pair<ADomain, std::vector<T>>>
137 static findHeaviestLeafRepeated(AItemInDomainMeasure& weightItemInDomain,
138 int maxLevel,
139 const TrackingUtilities::Weight minWeight = NAN)
140 {
141 auto skipLowWeightNode = [minWeight](const Node * node) {
142 return not(node->getWeight() >= minWeight);
143 };
144 return findHeaviestLeafRepeated(weightItemInDomain, maxLevel, skipLowWeightNode);
145 }
146
154 template <class AItemInDomainMeasure, class ASkipNodePredicate>
155 std::vector<std::pair<ADomain, std::vector<T>>>
156 findHeaviestLeafRepeated(AItemInDomainMeasure& weightItemInDomain,
157 int maxLevel,
158 ASkipNodePredicate& skipNode)
159 {
160 std::vector<std::pair<ADomain, std::vector<T> > > found;
161 Node* node = findHeaviestLeaf(weightItemInDomain, maxLevel, skipNode);
162 while (node) {
163 const ADomain* domain = node;
164 found.emplace_back(*domain, std::vector<T>(node->begin(), node->end()));
165 for (WithSharedMark<T>& markableItem : *node) {
166 markableItem.mark();
167 }
168 node = findHeaviestLeaf(weightItemInDomain, maxLevel, skipNode);
169 }
170 return found;
171 }
172
178 template <class AItemInDomainMeasure, class ASkipNodePredicate>
179 std::unique_ptr<std::pair<ADomain, std::vector<T>>>
180 findHeaviestLeafSingle(AItemInDomainMeasure& weightItemInDomain,
181 int maxLevel,
182 ASkipNodePredicate& skipNode)
183 {
184 using Result = std::pair<ADomain, std::vector<T> >;
185 std::unique_ptr<Result> found = nullptr;
186 Node* node = findHeaviestLeaf(weightItemInDomain, maxLevel, skipNode);
187 if (node) {
188 const ADomain* domain = node;
189 found.reset(new Result(*domain, std::vector<T>(node->begin(), node->end())));
190 for (WithSharedMark<T>& markableItem : *node) {
191 markableItem.mark();
192 }
193 }
194 return found;
195 }
196
202 template <class AItemInDomainMeasure, class ASkipNodePredicate>
203 Node* findHeaviestLeaf(AItemInDomainMeasure& weightItemInDomain,
204 int maxLevel,
205 ASkipNodePredicate& skipNode)
206 {
207 Node* heaviestNode = nullptr;
208 TrackingUtilities::Weight heighestWeigth = NAN;
209 auto isLeaf = [&heaviestNode, &heighestWeigth, maxLevel, &skipNode](Node * node) {
210 // Skip the expansion and the filling of the children
211 if (skipNode(node)) {
212 return true;
213 }
214
215 TrackingUtilities::Weight nodeWeight = node->getWeight();
216 // Skip the expansion and filling of the children if the node has not enough weight
217 if (not std::isnan(heighestWeigth) and not(nodeWeight > heighestWeigth)) {
218 return true;
219 }
220
221 // Node is a leaf at the maximum level and is heavier than everything seen before.
222 // Save its content
223 // Do not walk children
224 if (node->getLevel() >= maxLevel) {
225 heaviestNode = node;
226 heighestWeigth = nodeWeight;
227 return true;
228 }
229 return false;
230 };
231 // The isLeaf predicate does not mark any items.
232 const bool isLeafMarksItems = false;
233 fillWalk(weightItemInDomain, isLeaf, isLeafMarksItems);
234 return heaviestNode;
235 }
236
237 public:
242 template<class AItemInDomainMeasure, class AIsLeafPredicate>
243 void fillWalk(AItemInDomainMeasure& weightItemInDomain,
244 AIsLeafPredicate& isLeaf,
245 bool isLeafMarksItems = true)
246 {
247 auto walker = [&weightItemInDomain, &isLeaf](Node * node) {
248 // Check if node is a leaf
249 // Do not create children in this case
250 if (isLeaf(node)) {
251 // Do not walk children.
252 return false;
253 }
254
255 // Node is not a leaf.
256 // Check if it has children.
257 // If children have not been created, create and fill them.
258 typename Node::Children* children = node->getChildren();
259 if (not children) {
260 node->createChildren();
261 children = node->getChildren();
262 if constexpr(std::is_invocable_v<AItemInDomainMeasure&, const T&, Node*>) {
263 // Weighting function does not modify the item: fill all children in a single pass
264 // over the parent items without copies. Each child receives the items in the same order.
265 for (const WithSharedMark<T>& markableItem : *node) {
266 // Weighting function should not see the mark, but only the item itself.
267 const T& item(markableItem);
268 for (Node& childNode : *children) {
269 const TrackingUtilities::Weight weight = weightItemInDomain(item, &childNode);
270 if (not std::isnan(weight)) {
271 childNode.insert(markableItem, weight);
272 }
273 }
274 }
275 } else {
276 for (Node& childNode : *children) {
277 assert(childNode.getChildren() == nullptr);
278 assert(childNode.size() == 0);
279 auto measure =
280 // cppcheck-suppress constParameterReference ; the item is unwrapped as a non-const reference below
281 [&childNode, &weightItemInDomain](WithSharedMark<T>& markableItem) -> TrackingUtilities::Weight {
282 // Weighting function should not see the mark, but only the item itself.
283 T & item(markableItem);
284 return weightItemInDomain(item, &childNode);
285 };
286 childNode.insert(*node, measure);
287 }
288 }
289 }
290 // Continue to walk the children.
291 return true;
292 };
293 walkHeighWeightFirst(walker, isLeafMarksItems);
294 }
295
301 template<class ATreeWalker>
302 void walkHeighWeightFirst(ATreeWalker& walker, bool walkerMarksItems = true)
303 {
304 if (not walkerMarksItems and std::find(m_marks.begin(), m_marks.end(), true) == m_marks.end()) {
305 auto unmarkedPriority = [](Node * node) -> float {
306 return node->getWeight();
307 };
308 this->walk(walker, unmarkedPriority);
309 return;
310 }
311
312 auto priority = [](Node * node) -> float {
314 auto isMarked = [](const WithSharedMark<T>& markableItem) -> bool {
315 return markableItem.isMarked();
316 };
317 node->eraseIf(isMarked);
318 return node->getWeight();
319 };
320
321 this->walk(walker, priority);
322 }
323
325 // cppcheck-suppress duplInheritedMember ; intentionally hides the base class member, which it extends and then calls
326 void fell()
327 {
328 this->getTopNode().clear();
329 m_marks.clear();
330 Super::fell();
331 }
332
334 // cppcheck-suppress duplInheritedMember ; intentionally hides the base class member, which it extends and then calls
335 void raze()
336 {
337 this->fell();
338 Super::raze();
339 m_marks.shrink_to_fit();
340 }
341
342 private:
344 std::deque<bool> m_marks;
345 // Note: Have to use a deque here because std::vector<bool> is special
346 // std::vector<bool> m_marks;
347 };
348 }
350}
DynTree(const Properties &properties, const SubPropertiesFactory &subPropertiesFactory=SubPropertiesFactory())
Definition DynTree.h:221
Dynamic tree structure with weighted items in each node which are markable through out the tree.
void fell()
Fell to tree meaning deleting all child nodes from the tree. Keeps the top node.
void seed(const Ts &items)
Take the item set and insert them into the top node of the hough space.
Node * findHeaviestLeaf(AItemInDomainMeasure &weightItemInDomain, int maxLevel, ASkipNodePredicate &skipNode)
Go through all children until the maxLevel is reached and find the leaf with the highest weight.
std::vector< std::pair< ADomain, std::vector< T > > > findLeavesDisjoint(AItemInDomainMeasure &weightItemInDomain, int maxLevel, ASkipNodePredicate &skipNode)
Find all children node at maximum level and add them to the result list. Skip nodes if skipNode retur...
void raze()
Like fell but also releases all memory the tree has acquired during long executions.
static std::vector< std::pair< ADomain, std::vector< T > > > findHeaviestLeafRepeated(AItemInDomainMeasure &weightItemInDomain, int maxLevel, const TrackingUtilities::Weight minWeight=NAN)
Go through all children until maxLevel is reached and find the heaviest leaves.
void fillWalk(AItemInDomainMeasure &weightItemInDomain, AIsLeafPredicate &isLeaf, bool isLeafMarksItems=true)
Walk through the children and fill them if necessary until isLeaf returns true.
std::unique_ptr< std::pair< ADomain, std::vector< T > > > findHeaviestLeafSingle(AItemInDomainMeasure &weightItemInDomain, int maxLevel, ASkipNodePredicate &skipNode)
Go through all children until the maxLevel is reached and find the leaf with the highest weight.
void walkHeighWeightFirst(ATreeWalker &walker, bool walkerMarksItems=true)
Walk the tree investigating the heaviest children with priority.
std::vector< std::pair< ADomain, std::vector< T > > > findHeavyLeavesDisjoint(AItemInDomainMeasure &weightItemInDomain, int maxLevel, double minWeight)
Find all children node at maximum level and add them to the result list. Skip nodes if their weight i...
std::vector< std::pair< ADomain, std::vector< T > > > findHeaviestLeafRepeated(AItemInDomainMeasure &weightItemInDomain, int maxLevel, ASkipNodePredicate &skipNode)
Go through all children until maxLevel is reached and find the heaviest leaves.
WeightedParititioningDynTree< WithSharedMark< AItemPtr >, HoughBox, BoxDivision > Super
DynTree< WithWeightedItems< ADomain, T >, ADomainDivsion > Super
Type of the base class.
WeightedParititioningDynTree(ADomain topDomain, ADomainDivsion domainDivsion)
Constructor attaching a vector of the weighted items to the top most node domain.
Mixin class to attach a mark that is shared among many instances.
A mixin class to attach a set of weighted items to a class.
Abstract base class for different kinds of events.
STL namespace.