ATLAS Offline Software
Loading...
Searching...
No Matches
SegmentEdgeInferenceAlg.cxx
Go to the documentation of this file.
2#include "InferenceUtils.h"
7#include "Acts/Utilities/Helpers.hpp"
8#include <algorithm>
9#include <cmath>
10#include <cstdint>
11#include <limits>
12#include <memory>
13#include <numeric>
14#include <unordered_map>
15#include <unordered_set>
16#include <utility>
17#include <vector>
18
19namespace MuonML {
20
21namespace {
22
23std::uint64_t undirectedPairKey(std::size_t first, std::size_t second) {
24 if (first > second) std::swap(first, second);
25 return (static_cast<std::uint64_t>(first) << 32) |
26 static_cast<std::uint64_t>(second);
27}
28
29class DisjointSet {
30public:
31 explicit DisjointSet(std::size_t size) : m_parent(size), m_rank(size, 0) {
32 std::iota(m_parent.begin(), m_parent.end(), 0);
33 }
34
35 std::size_t find(std::size_t node) {
36 if (m_parent[node] != node) m_parent[node] = find(m_parent[node]);
37 return m_parent[node];
38 }
39
40 void unite(std::size_t first, std::size_t second) {
41 first = find(first);
42 second = find(second);
43 if (first == second) return;
44 if (m_rank[first] < m_rank[second]) std::swap(first, second);
45 m_parent[second] = first;
46 if (m_rank[first] == m_rank[second]) ++m_rank[first];
47 }
48
49private:
50 std::vector<std::size_t> m_parent;
51 std::vector<unsigned char> m_rank;
52};
53
54} // namespace
55
57 ATH_CHECK(m_segmentKey.initialize());
58 ATH_CHECK(m_pairGateDecorKey.initialize());
60 ATH_CHECK(m_edgeClassifier.retrieve());
61 return StatusCode::SUCCESS;
62}
63
64StatusCode SegmentEdgeInferenceAlg::execute(const EventContext& ctx) const {
65 const xAOD::MuonSegmentContainer* segments{};
66 ATH_CHECK(SG::get(segments, m_segmentKey, ctx));
67 ATH_MSG_DEBUG("Event " << ctx.eventID().event_number()
68 << ": input segments in '" << m_segmentKey.key()
69 << "' = " << segments->size());
70
71 SegmentEdgeGraph graph{};
72 std::vector<SegmentEdgeScore> scores{};
73 ATH_CHECK(m_edgeClassifier->buildGraph(ctx, *segments, graph));
74 ATH_MSG_DEBUG("Event " << ctx.eventID().event_number()
75 << ": built graph with nodes=" << graph.nNodes
76 << ", edges=" << graph.nEdges);
77
78 ATH_CHECK(m_edgeClassifier->classifyEdges(ctx, graph, scores));
79 if (!scores.empty()) {
80 float minProb = std::numeric_limits<float>::max();
81 float maxProb = std::numeric_limits<float>::lowest();
82 for (const SegmentEdgeScore& score : scores) {
83 minProb = std::min(minProb, score.probability);
84 maxProb = std::max(maxProb, score.probability);
85 }
86 ATH_MSG_DEBUG("Event " << ctx.eventID().event_number()
87 << ": edge scores=" << scores.size()
88 << ", prob range=[" << minProb << ", " << maxProb << "]");
89 } else {
90 ATH_MSG_DEBUG("Event " << ctx.eventID().event_number()
91 << ": no edge scores produced");
92 }
93
94 // Symmetrise directed model outputs once. Everything downstream uses the
95 // score of an undirected segment association.
96 std::unordered_map<std::uint64_t, float> pairProbability;
97 pairProbability.reserve(scores.size());
98 for (const SegmentEdgeScore& score : scores) {
99 if (score.src >= graph.nNodes || score.dst >= graph.nNodes ||
100 score.src == score.dst) {
101 continue;
102 }
103 const std::uint64_t key = undirectedPairKey(score.src, score.dst);
104 auto [it, inserted] = pairProbability.emplace(key, score.probability);
105 if (!inserted) it->second = std::max(it->second, score.probability);
106 }
107
108 // Form a sparse, score-ranked topology. Mutual top-K associations retain
109 // locally consistent paths while suppressing one-sided bridges.
110 using WeightedEdge = std::pair<std::uint64_t, float>;
111 const auto betterWeightedEdge = [](const WeightedEdge& first,
112 const WeightedEdge& second) {
113 const int probabilityOrder = InferenceUtils::compareFloatDescending(
114 first.second, second.second);
115 if (probabilityOrder != 0) {
116 return probabilityOrder < 0;
117 }
118 return first.first < second.first;
119 };
120 std::vector<std::vector<WeightedEdge>> edgesByNode(graph.nNodes);
121 std::size_t thresholdPairs = 0;
122 for (const auto& [key, probability] : pairProbability) {
123 if (probability < m_pairGateThreshold.value()) continue;
124 const std::size_t first = static_cast<std::size_t>(key >> 32);
125 const std::size_t second = static_cast<std::size_t>(key & 0xffffffffu);
126 if (first >= graph.nNodes || second >= graph.nNodes) continue;
127 edgesByNode[first].emplace_back(key, probability);
128 edgesByNode[second].emplace_back(key, probability);
129 ++thresholdPairs;
130 }
131
132 std::size_t thresholdedNodes = 0;
133 for (const std::vector<WeightedEdge>& nodeEdges : edgesByNode) {
134 thresholdedNodes += !nodeEdges.empty();
135 }
136
137 std::unordered_map<std::uint64_t, unsigned char> nominations;
138 nominations.reserve(thresholdPairs);
139 for (std::vector<WeightedEdge>& nodeEdges : edgesByNode) {
140 std::sort(nodeEdges.begin(), nodeEdges.end(), betterWeightedEdge);
141 if (m_maxEdgesPerNode.value() != 0 &&
142 nodeEdges.size() > m_maxEdgesPerNode.value()) {
143 nodeEdges.resize(m_maxEdgesPerNode.value());
144 }
145 for (const WeightedEdge& edge : nodeEdges) {
146 ++nominations[edge.first];
147 }
148 }
149
150 std::size_t mutualTopKPairs = 0;
151 std::size_t oneSidedTopKPairs = 0;
152 for (const auto& [_, count] : nominations) {
153 if (count == 2) {
154 ++mutualTopKPairs;
155 } else {
156 ++oneSidedTopKPairs;
157 }
158 }
159
160 std::unordered_set<std::uint64_t> selectedPairKeys;
161 selectedPairKeys.reserve(thresholdPairs);
162 if (m_useDegreeCappedComponents.value()) {
163 // Apply a score-ordered global degree cap when explicitly requested.
164 std::vector<WeightedEdge> acceptedPairs;
165 acceptedPairs.reserve(thresholdPairs);
166 for (const auto& [key, probability] : pairProbability) {
167 if (probability < m_pairGateThreshold.value()) continue;
168 acceptedPairs.emplace_back(key, probability);
169 }
170 std::sort(acceptedPairs.begin(), acceptedPairs.end(), betterWeightedEdge);
171
172 const unsigned int maxDegree = m_maxEdgesPerNode.value();
173 std::vector<unsigned int> degree(graph.nNodes, 0);
174 for (const WeightedEdge& edge : acceptedPairs) {
175 const std::size_t first = static_cast<std::size_t>(edge.first >> 32);
176 const std::size_t second = static_cast<std::size_t>(edge.first & 0xffffffffu);
177 if (maxDegree != 0 &&
178 (degree[first] >= maxDegree || degree[second] >= maxDegree)) {
179 continue;
180 }
181 selectedPairKeys.insert(edge.first);
182 ++degree[first];
183 ++degree[second];
184 }
185 } else {
186 for (const auto& [key, count] : nominations) {
187 if (m_requireMutualTopKEdges.value() && count != 2) continue;
188 selectedPairKeys.insert(key);
189 }
190 }
191
192 // A mutual top-K selection can leave an endpoint without an association.
193 // Add its best thresholded edge, at most once per orphaned endpoint.
194 std::size_t orphanRecoveryPairs = 0;
195 if (!m_useDegreeCappedComponents.value() &&
196 m_requireMutualTopKEdges.value() &&
197 m_recoverOrphanNodes.value()) {
198 std::vector<unsigned char> selectedNode(graph.nNodes, 0);
199 for (const std::uint64_t key : selectedPairKeys) {
200 const std::size_t first = static_cast<std::size_t>(key >> 32);
201 const std::size_t second =
202 static_cast<std::size_t>(key & 0xffffffffu);
203 if (first < graph.nNodes) selectedNode[first] = 1;
204 if (second < graph.nNodes) selectedNode[second] = 1;
205 }
206 for (std::size_t node = 0; node < graph.nNodes; ++node) {
207 if (selectedNode[node] || edgesByNode[node].empty()) continue;
208 const std::uint64_t key = edgesByNode[node].front().first;
209 const std::size_t first = static_cast<std::size_t>(key >> 32);
210 const std::size_t second =
211 static_cast<std::size_t>(key & 0xffffffffu);
212 if (first >= graph.nNodes || second >= graph.nNodes) continue;
213 if (selectedPairKeys.insert(key).second) ++orphanRecoveryPairs;
214 selectedNode[first] = 1;
215 selectedNode[second] = 1;
216 }
217 }
218
219 std::vector<std::uint64_t> selectedPairs{selectedPairKeys.begin(),
220 selectedPairKeys.end()};
221 std::sort(selectedPairs.begin(), selectedPairs.end());
222
223 DisjointSet components{graph.nNodes};
224 std::vector<bool> activeNode(graph.nNodes, false);
225 for (const std::uint64_t key : selectedPairs) {
226 const std::size_t first = static_cast<std::size_t>(key >> 32);
227 const std::size_t second = static_cast<std::size_t>(key & 0xffffffffu);
228 components.unite(first, second);
229 activeNode[first] = true;
230 activeNode[second] = true;
231 }
232
233 std::unordered_map<std::size_t, std::vector<std::size_t>> byRoot;
234 byRoot.reserve(graph.nNodes);
235 for (std::size_t node = 0; node < graph.nNodes; ++node) {
236 if (activeNode[node]) byRoot[components.find(node)].push_back(node);
237 }
238
239 // Deterministic component IDs make debugging and validation reproducible.
240 std::vector<std::vector<std::size_t>> componentNodes;
241 componentNodes.reserve(byRoot.size());
242 for (auto& [_, nodes] : byRoot) {
243 std::sort(nodes.begin(), nodes.end());
244 componentNodes.push_back(std::move(nodes));
245 }
246 std::ranges::sort(componentNodes,
247 [](const auto& first, const auto& second) {
248 return first.front() < second.front();
249 });
250
252 decor{m_pairGateDecorKey, ctx};
253
254 if (!m_filteredSegmentKey.empty()) {
255 auto connectedSegments =
256 std::make_unique<ConstDataVector<xAOD::MuonSegmentContainer>>(
258 connectedSegments->reserve(graph.nNodes);
259 for (std::size_t node = 0; node < graph.nNodes; ++node) {
260 if (!activeNode[node] || !graph.segments[node]) continue;
261 connectedSegments->push_back(graph.segments[node]);
262 }
263
264 const std::size_t nConnectedSegments = connectedSegments->size();
267 ATH_CHECK(connectedHandle.record(std::move(connectedSegments)));
268 ATH_MSG_DEBUG("Event " << ctx.eventID().event_number()
269 << ": wrote " << nConnectedSegments
270 << " ML-connected segment(s) to '"
271 << m_filteredSegmentKey.key() << "'");
272 }
273
274 std::size_t topologyNodes = 0;
275 std::size_t retainedNodes = 0;
276 std::size_t chamberSuppressedNodes = 0;
277 std::size_t rejectedComponents = 0;
278 std::size_t componentsKept = 0;
279 std::size_t anchors = 0;
280 std::size_t nodesRejectedByMinComponent = 0;
281 unsigned nextComponentId = 1;
282 const InferenceUtils::SegmentQualityOrder betterSegment{};
283
284 const auto isBetterNode = [&](std::size_t candidate, std::size_t incumbent) {
285 if (betterSegment(graph.segments[candidate], graph.segments[incumbent])) {
286 return true;
287 }
288 if (betterSegment(graph.segments[incumbent], graph.segments[candidate])) {
289 return false;
290 }
291 return candidate < incumbent;
292 };
293
294 for (const std::vector<std::size_t>& rawNodes : componentNodes) {
295 topologyNodes += rawNodes.size();
296 std::vector<std::size_t> retained = rawNodes;
297
298 // The seeder resolves same-chamber alternatives while building a seed.
299 // Retaining only the highest-ranked representative is therefore optional.
300 if (m_keepBestSegmentPerChamber.value()) {
301 std::unordered_map<int, std::size_t> bestByChamber;
302 bestByChamber.reserve(rawNodes.size());
303 for (const std::size_t node : rawNodes) {
304 const int chamber =
305 static_cast<int>(graph.segments[node]->chamberIndex());
306 const auto found = bestByChamber.find(chamber);
307 if (found == bestByChamber.end() || isBetterNode(node, found->second)) {
308 bestByChamber[chamber] = node;
309 }
310 }
311 retained.clear();
312 retained.reserve(bestByChamber.size());
313 for (const auto& [_, node] : bestByChamber) retained.push_back(node);
314 std::sort(retained.begin(), retained.end());
315 chamberSuppressedNodes += rawNodes.size() - retained.size();
316 }
317
318 if (retained.size() < m_minSegmentsPerComponent.value()) {
319 nodesRejectedByMinComponent += retained.size();
320 ++rejectedComponents;
321 continue;
322 }
323
324 // Only ranked component members launch seeds. The edge score therefore
325 // reduces seed attempts directly rather than serving only as a label.
326 std::vector<std::size_t> rankedNodes{retained};
327 std::ranges::sort(rankedNodes, isBetterNode);
328 const std::size_t nAnchors = m_seedAnchorsPerComponent.value() == 0
329 ? rankedNodes.size()
330 : std::min<std::size_t>(m_seedAnchorsPerComponent.value(),
331 rankedNodes.size());
332 if (nAnchors == 0) {
333 ++rejectedComponents;
334 continue;
335 }
336 rankedNodes.resize(nAnchors);
337 std::ranges::sort(rankedNodes);
338
339 const unsigned componentId = nextComponentId++;
340 for (const std::size_t node : retained) {
341 const bool isAnchor = Acts::rangeContainsValue(rankedNodes, node);
342 decor(*graph.segments[node]) = {
343 componentId, static_cast<unsigned int>(isAnchor)};
344 }
345 retainedNodes += retained.size();
346 anchors += nAnchors;
347 ++componentsKept;
348 }
349
350 ATH_MSG_DEBUG("Event " << ctx.eventID().event_number()
351 << ": ML components graphNodes=" << graph.nNodes
352 << ", thresholdedNodes=" << thresholdedNodes
353 << ", thresholdPairs=" << thresholdPairs
354 << ", mutualTopKPairs=" << mutualTopKPairs
355 << ", oneSidedTopKPairs=" << oneSidedTopKPairs
356 << ", selectedPairs=" << selectedPairKeys.size()
357 << ", components=" << componentsKept
358 << ", topologyNodes=" << topologyNodes
359 << ", retainedNodes=" << retainedNodes
360 << ", chamberSuppressedNodes=" << chamberSuppressedNodes
361 << ", nodesRejectedByMinComponent=" << nodesRejectedByMinComponent
362 << ", keepBestSegmentPerChamber=" << m_keepBestSegmentPerChamber.value()
363 << ", seedAnchors=" << anchors
364 << ", rejectedComponents=" << rejectedComponents
365 << ", threshold=" << m_pairGateThreshold.value()
366 << ", maxEdgesPerNode=" << m_maxEdgesPerNode.value()
367 << ", orphanRecoveryPairs=" << orphanRecoveryPairs
368 << ", mutualTopK=" << m_requireMutualTopKEdges.value()
369 << ", degreeCapped=" << m_useDegreeCappedComponents.value());
370 return StatusCode::SUCCESS;
371}
372
373} // namespace MuonML
#define ATH_CHECK
Evaluate an expression and check for errors.
#define ATH_MSG_DEBUG(x,...)
DataVector adapter that acts like it holds const pointers.
Handle class for reading from StoreGate.
Handle class for recording to StoreGate.
size_t size() const
Number of registered mappings.
static const Attributes_t empty
size_type size() const noexcept
Returns the number of elements in the collection.
Gaudi::Property< bool > m_keepBestSegmentPerChamber
SG::WriteDecorHandleKey< xAOD::MuonSegmentContainer > m_pairGateDecorKey
Per-segment payload consumed by MlMsTrackSeeder: [componentId, isSeedAnchor] Empty means that the seg...
Gaudi::Property< bool > m_useDegreeCappedComponents
ToolHandle< ISegmentEdgeClassifierTool > m_edgeClassifier
Gaudi::Property< unsigned int > m_maxEdgesPerNode
Gaudi::Property< bool > m_requireMutualTopKEdges
StatusCode execute(const EventContext &ctx) const override
Gaudi::Property< unsigned int > m_minSegmentsPerComponent
SG::WriteHandleKey< ConstDataVector< xAOD::MuonSegmentContainer > > m_filteredSegmentKey
Gaudi::Property< bool > m_recoverOrphanNodes
SG::ReadHandleKey< xAOD::MuonSegmentContainer > m_segmentKey
Gaudi::Property< unsigned int > m_seedAnchorsPerComponent
Gaudi::Property< float > m_pairGateThreshold
StatusCode record(std::unique_ptr< T > data)
Record a const object to the store.
Definition node.h:24
Auxiliary class to instantiate WriteDecorHandles.The handles can be created in an empty state.
std::string find(const std::string &s)
return a remapped string
Definition hcg.cxx:140
int count(std::string s, const std::string &regx)
count how many occurances of a regx are in a string
Definition hcg.cxx:148
bool first
Definition DeMoScan.py:534
int compareFloatDescending(float first, float second)
Three-way descending comparison which also orders NaN last.
@ VIEW_ELEMENTS
this data object is a view, it does not own its elmts
const T * get(const ReadCondHandleKey< T > &key, const EventContext &ctx)
Convenience function to retrieve an object given a ReadCondHandleKey.
void sort(typename DataModel_detail::iterator< DVL > beg, typename DataModel_detail::iterator< DVL > end)
Specialization of sort for DataVector/List.
void swap(ElementLinkVector< DOBJ > &lhs, ElementLinkVector< DOBJ > &rhs)
MuonSegmentContainer_v1 MuonSegmentContainer
Definition of the current "MuonSegment container version".
Common quality ordering for segment representatives.
std::vector< const xAOD::MuonSegment_v1 * > segments