2222import time
2323from collections import defaultdict
2424from collections .abc import Iterable , Mapping , Sequence
25- from itertools import groupby
25+ from itertools import groupby , chain
2626from multiprocessing import Pool
2727from operator import index
2828from pathlib import Path
@@ -80,6 +80,7 @@ class DataCollection:
8080 def __init__ (
8181 self , files , sources_data = None , train_ids = None , aliases = None ,
8282 ctx_closes = False , * , inc_suspect_trains = True , is_single_run = False ,
83+ alias_files = None
8384 ):
8485 self .files = list (files )
8586 self .ctx_closes = ctx_closes
@@ -117,6 +118,9 @@ def __init__(
117118 }
118119 self ._sources_data = sources_data
119120
121+ # Note that _alias_files is only for tracking where the aliases came
122+ # from, the actual aliases are stored in _aliases.
123+ self ._alias_files = [] if alias_files is None else alias_files
120124 self ._aliases = aliases or {}
121125 self .alias = AliasIndexer (self )
122126
@@ -629,6 +633,10 @@ def _merge_aliases(self, alias_dicts):
629633
630634 return new_aliases
631635
636+ def _merge_alias_files (self , * alias_files ):
637+ all_files = chain .from_iterable (alias_files )
638+ return sorted (set (all_files ))
639+
632640 def union (self , * others ):
633641 """Join the data in this collection with one or more others.
634642
@@ -653,6 +661,8 @@ def union(self, *others):
653661
654662 aliases = self ._merge_aliases (
655663 [self ._aliases ] + [dc ._aliases for dc in others ])
664+ alias_files = self ._merge_alias_files (self ._alias_files ,
665+ * [dc ._alias_files for dc in others ])
656666
657667 train_ids = sorted (set ().union (* [sd .train_ids for sd in sources_data .values ()]))
658668 # Update the internal list of train IDs for the sources
@@ -664,7 +674,7 @@ def union(self, *others):
664674 return DataCollection (
665675 files , sources_data = sources_data , train_ids = train_ids ,
666676 aliases = aliases , inc_suspect_trains = self .inc_suspect_trains ,
667- is_single_run = same_run (self , * others ),
677+ is_single_run = same_run (self , * others ), alias_files = alias_files
668678 )
669679
670680 def __or__ (self , other ):
@@ -677,6 +687,7 @@ def _parse_aliases(self, alias_defs):
677687 """Parse alias definitions into alias dictionaries."""
678688
679689 alias_dicts = []
690+ alias_files = []
680691
681692 def is_valid_alias (k , v ):
682693 return (isinstance (k , str ) and (
@@ -693,10 +704,11 @@ def is_valid_alias(k, v):
693704 alias_dicts .append (alias_def )
694705 elif isinstance (alias_def , (str , os .PathLike )):
695706 # From a file.
707+ alias_files .append (Path (alias_def ))
696708 alias_dicts .append (
697709 self ._load_aliases_from_file (Path (alias_def )))
698710
699- return alias_dicts
711+ return alias_dicts , alias_files
700712
701713 def _load_aliases_from_file (self , aliases_path ):
702714 """Load alias definitions from file."""
@@ -786,14 +798,15 @@ def with_aliases(self, *alias_defs):
786798 """
787799
788800 # Check for conflicts within these definitions
789- new_aliases = self ._merge_aliases (
790- [self ._aliases ] + self ._parse_aliases (alias_defs ))
801+ new_aliases , new_alias_files = self ._parse_aliases (alias_defs )
802+ new_aliases = self ._merge_aliases ([self ._aliases ] + new_aliases )
803+ alias_files = self ._merge_alias_files (self ._alias_files , new_alias_files )
791804
792805 return DataCollection (
793806 self .files , sources_data = self ._sources_data ,
794807 train_ids = self .train_ids , aliases = new_aliases ,
795808 inc_suspect_trains = self .inc_suspect_trains ,
796- is_single_run = self .is_single_run
809+ is_single_run = self .is_single_run , alias_files = alias_files
797810 )
798811
799812 def only_aliases (self , * alias_defs , strict = False , require_all = False ):
@@ -818,8 +831,9 @@ def only_aliases(self, *alias_defs, strict=False, require_all=False):
818831 """
819832
820833 # Create new aliases.
821- aliases = self ._merge_aliases (
822- [self ._aliases ] + self ._parse_aliases (alias_defs ))
834+ new_aliases , new_alias_files = self ._parse_aliases (alias_defs )
835+ aliases = self ._merge_aliases ([self ._aliases ] + new_aliases )
836+ alias_files = self ._merge_alias_files (self ._alias_files , new_alias_files )
823837
824838 # Set of sources aliased.
825839 aliased_sources = {literal for literal in aliases .values ()
@@ -865,6 +879,7 @@ def only_aliases(self, *alias_defs, strict=False, require_all=False):
865879 # Create a new DataCollection from selecting and add the aliases.
866880 new_data = self .select (selection , require_all = require_all )
867881 new_data ._aliases = aliases
882+ new_data ._alias_files = alias_files
868883
869884 return new_data
870885
@@ -1092,7 +1107,7 @@ def select(self, seln_or_source_glob, key_glob='*', require_all=False,
10921107 return DataCollection (
10931108 files , sources_data , train_ids = train_ids , aliases = self ._aliases ,
10941109 inc_suspect_trains = self .inc_suspect_trains ,
1095- is_single_run = self .is_single_run
1110+ is_single_run = self .is_single_run , alias_files = self . _alias_files
10961111 )
10971112
10981113 def deselect (self , seln_or_source_glob , key_glob = '*' ):
@@ -1129,7 +1144,7 @@ def deselect(self, seln_or_source_glob, key_glob='*'):
11291144 return DataCollection (
11301145 files , sources_data = sources_data , train_ids = self .train_ids ,
11311146 aliases = self ._aliases , inc_suspect_trains = self .inc_suspect_trains ,
1132- is_single_run = self .is_single_run ,
1147+ is_single_run = self .is_single_run , alias_files = self . _alias_files
11331148 )
11341149
11351150 def select_trains (self , train_range ):
@@ -1164,7 +1179,7 @@ def select_trains(self, train_range):
11641179 return DataCollection (
11651180 files , sources_data = sources_data , train_ids = new_train_ids ,
11661181 aliases = self ._aliases , inc_suspect_trains = self .inc_suspect_trains ,
1167- is_single_run = self .is_single_run ,
1182+ is_single_run = self .is_single_run , alias_files = self . _alias_files
11681183 )
11691184
11701185 def split_trains (self , parts = None , trains_per_part = None ):
@@ -1207,7 +1222,7 @@ def dict_zip(iter_d):
12071222 yield DataCollection (
12081223 files , sources_data = sources_data_part , train_ids = train_ids ,
12091224 aliases = self ._aliases , inc_suspect_trains = self .inc_suspect_trains ,
1210- is_single_run = self .is_single_run ,
1225+ is_single_run = self .is_single_run , alias_files = self . _alias_files
12111226 )
12121227
12131228 def _check_source_conflicts (self ):
0 commit comments