[BUG]: PatternDataGrabber issues with replacements #492
2 changed files with 15 additions and 1 deletions
1
docs/changes/newsfragments/492.bugfix
Normal file
1
docs/changes/newsfragments/492.bugfix
Normal file
|
|
@ -0,0 +1 @@
|
||||||
|
Fix pattern indexing logic for :meth:`.PatternDataGrabber.get_elements` by `Synchon Mandal`_
|
||||||
|
|
@ -449,6 +449,19 @@ class PatternDataGrabber(BaseDataGrabber, PatternValidationMixin):
|
||||||
|
|
||||||
return out
|
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:
|
def get_elements(self) -> Elements:
|
||||||
"""Implement fetching list of elements in the dataset.
|
"""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
|
# Iterate by number of replacements. At least one pattern must have
|
||||||
# all of them.
|
# all of them.
|
||||||
replacement_in_patterns = [
|
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
|
for t_type in self.types
|
||||||
]
|
]
|
||||||
order = np.argsort(replacement_in_patterns)
|
order = np.argsort(replacement_in_patterns)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue