[BUG]: PatternDataGrabber issues with replacements #492

Merged
synchon merged 3 commits from replacements-in-patterns into main 2026-03-26 14:21:33 +00:00
2 changed files with 15 additions and 1 deletions

View file

@ -0,0 +1 @@
Fix pattern indexing logic for :meth:`.PatternDataGrabber.get_elements` by `Synchon Mandal`_

View file

@ -449,6 +449,19 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
return out
def _count_replacements_in_pattern(self, dtype: str) -> int:
"""Count the replacements in pattern(s)."""
dtype_val = self.patterns[dtype]
# Convert non-list dtype vals
if not isinstance(dtype_val, list):
dtype_val = [dtype_val]
counts = []
for val in dtype_val:
counts.append(
np.sum([x in val.get("pattern") for x in self.replacements])
)
return np.max(counts).astype(int)
def get_elements(self) -> Elements:
"""Implement fetching list of elements in the dataset.
@ -467,7 +480,7 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
# Iterate by number of replacements. At least one pattern must have
# all of them.
replacement_in_patterns = [
np.sum([x in self.patterns[t_type] for x in self.replacements])
self._count_replacements_in_pattern(t_type)
for t_type in self.types
]
order = np.argsort(replacement_in_patterns)