2121from lavender_data .server .reader import (
2222 get_reader_instance ,
2323 GlobalSampleIndex ,
24+ JoinMethod ,
25+ InnerJoinSampleInsufficient ,
2426)
2527from lavender_data .server .registries import (
2628 PreprocessorRegistry ,
@@ -107,7 +109,13 @@ def _decollate(batch: dict) -> dict:
107109 return _batch
108110
109111
110- def _process_next_samples (params : ProcessNextSamplesParams ) -> dict :
112+ class NoSamplesFound (Exception ):
113+ pass
114+
115+
116+ def _process_next_samples (
117+ params : ProcessNextSamplesParams , join_method : JoinMethod = "left"
118+ ) -> dict :
111119 reader = get_reader_instance ()
112120
113121 current = params .current
@@ -118,7 +126,15 @@ def _process_next_samples(params: ProcessNextSamplesParams) -> dict:
118126 batch_size = params .batch_size
119127
120128 if samples is None :
121- samples = [reader .get_sample (i , join = "left" ) for i in global_sample_indices ]
129+ samples = []
130+ for i in global_sample_indices :
131+ try :
132+ samples .append (reader .get_sample (i , join_method ))
133+ except InnerJoinSampleInsufficient :
134+ pass
135+
136+ if len (samples ) == 0 :
137+ raise NoSamplesFound ()
122138
123139 batch = (
124140 CollaterRegistry .get (collater ["name" ]).collate (samples )
@@ -146,18 +162,24 @@ def _process_next_samples(params: ProcessNextSamplesParams) -> dict:
146162def process_next_samples (
147163 params : ProcessNextSamplesParams ,
148164 max_retry_count : int ,
165+ join_method : JoinMethod = "left" ,
149166) -> dict :
150- logger = get_logger (__name__ )
151-
152167 for i in range (max_retry_count + 1 ):
153168 try :
154- return _process_next_samples (params )
169+ return _process_next_samples (params , join_method )
170+ except NoSamplesFound as e :
171+ raise ProcessNextSamplesException (
172+ e = e ,
173+ current = params .current ,
174+ global_sample_indices = params .global_sample_indices ,
175+ )
155176 except Exception as e :
156177 error = ProcessNextSamplesException (
157178 e = e ,
158179 current = params .current ,
159180 global_sample_indices = params .global_sample_indices ,
160181 )
182+ logger = get_logger (__name__ )
161183 if i < max_retry_count :
162184 logger .warning (f"{ str (error )} , retrying... ({ i + 1 } /{ max_retry_count } )" )
163185 else :
0 commit comments