1515from collections .abc import Sequence
1616from typing import Any
1717
18+ from lhotse import CutSet
1819from lhotse .dataset import DynamicCutSampler
1920from lhotse .dataset .sampling .dynamic import DurationBatcher , Filter
2021from lhotse .lazy import get_graph_origin , resolve_iterator_source
2122
22- from lhotse import CutSet
23-
2423
25- def _select_best_fit_indices (
26- lengths : Sequence [int ], capacity : int , max_items : int | None = None
27- ) -> list [int ]:
24+ def _select_best_fit_indices (lengths : Sequence [int ], capacity : int , max_items : int | None = None ) -> list [int ]:
2825 """Select an exact best-fit subset, preferring earlier items on ties."""
2926 if capacity < 0 :
3027 raise ValueError (f"capacity must be non-negative (got { capacity } )" )
@@ -91,8 +88,7 @@ def __init__(self, *args, packing_buffer_size: int, **kwargs):
9188 super ().__init__ (* args , ** kwargs )
9289 if packing_buffer_size <= 0 :
9390 raise ValueError (
94- "shuffle_buffer_size must be a positive packing-buffer size "
95- f"(got { packing_buffer_size } )"
91+ "shuffle_buffer_size must be a positive packing-buffer size " f"(got { packing_buffer_size } )"
9692 )
9793 self .packing_buffer_size = packing_buffer_size
9894 self ._source_exhausted = False
@@ -108,32 +104,20 @@ def _detuplify(examples):
108104
109105 @staticmethod
110106 def _measured_example (example_or_tuple ):
111- return (
112- example_or_tuple [0 ]
113- if isinstance (example_or_tuple , tuple )
114- else example_or_tuple
115- )
107+ return example_or_tuple [0 ] if isinstance (example_or_tuple , tuple ) else example_or_tuple
116108
117109 def _fill_packing_buffer (self ) -> None :
118- while (
119- len (self .reuse_cuts_buffer ) < self .packing_buffer_size
120- and not self ._source_exhausted
121- ):
110+ while len (self .reuse_cuts_buffer ) < self .packing_buffer_size and not self ._source_exhausted :
122111 try :
123112 self .reuse_cuts_buffer .append (next (self .cuts_iter ))
124113 except StopIteration :
125114 self ._source_exhausted = True
126115
127116 def _measure_integer_length (self , example_or_tuple ) -> int :
128- length = self .constraint .measure_length (
129- self ._measured_example (example_or_tuple )
130- )
117+ length = self .constraint .measure_length (self ._measured_example (example_or_tuple ))
131118 integer_length = int (length )
132119 if integer_length != length :
133- raise ValueError (
134- "Packed sequence sampling requires integer token lengths, "
135- f"but measured { length !r} ."
136- )
120+ raise ValueError ("Packed sequence sampling requires integer token lengths, " f"but measured { length !r} ." )
137121 return integer_length
138122
139123 def _limits (self ) -> tuple [int , int | None ]:
@@ -144,18 +128,14 @@ def _limits(self) -> tuple[int, int | None]:
144128 max_tokens = getattr (internal , "max_tokens" , None )
145129 max_examples = getattr (internal , "max_examples" , max_examples )
146130 if max_tokens is None :
147- raise ValueError (
148- "Packed sequence sampling requires batch_tokens to define the exact token cap."
149- )
131+ raise ValueError ("Packed sequence sampling requires batch_tokens to define the exact token cap." )
150132 max_tokens = int (max_tokens )
151133 if max_tokens <= 0 :
152134 raise ValueError (f"batch_tokens must be positive (got { max_tokens } )" )
153135 if max_examples is not None :
154136 max_examples = int (max_examples )
155137 if max_examples <= 0 :
156- raise ValueError (
157- f"batch_size must be positive or null (got { max_examples } )"
158- )
138+ raise ValueError (f"batch_size must be positive or null (got { max_examples } )" )
159139 return max_tokens , max_examples
160140
161141 def _discard (self , examples ) -> None :
@@ -182,35 +162,21 @@ def _collect_batch(self):
182162 )
183163
184164 remaining_items = None if max_examples is None else max_examples - 1
185- tail_indices = _select_best_fit_indices (
186- lengths [1 :], max_tokens - anchor_length , max_items = remaining_items
187- )
165+ tail_indices = _select_best_fit_indices (lengths [1 :], max_tokens - anchor_length , max_items = remaining_items )
188166 selected_indices = {0 , * (index + 1 for index in tail_indices )}
189- examples = [
190- example for index , example in enumerate (pool ) if index in selected_indices
191- ]
192- deferred = [
193- example
194- for index , example in enumerate (pool )
195- if index not in selected_indices
196- ]
167+ examples = [example for index , example in enumerate (pool ) if index in selected_indices ]
168+ deferred = [example for index , example in enumerate (pool ) if index not in selected_indices ]
197169 self .reuse_cuts_buffer .clear ()
198170 self .reuse_cuts_buffer .extend (deferred )
199171
200172 self .constraint .reset ()
201173 for example in examples :
202174 self .constraint .add (self ._measured_example (example ))
203175 if self .constraint .exceeded ():
204- raise AssertionError (
205- "Best-fit packed batch exceeded its configured constraint."
206- )
176+ raise AssertionError ("Best-fit packed batch exceeded its configured constraint." )
207177
208178 is_final_batch = self ._source_exhausted and not self .reuse_cuts_buffer
209- if (
210- is_final_batch
211- and self .drop_last
212- and not self .constraint .close_to_exceeding ()
213- ):
179+ if is_final_batch and self .drop_last and not self .constraint .close_to_exceeding ():
214180 self ._discard (examples )
215181 raise StopIteration ()
216182
@@ -234,8 +200,7 @@ def __init__(
234200 ):
235201 if shuffle_buffer_size is None or shuffle_buffer_size <= 0 :
236202 raise ValueError (
237- "shuffle_buffer_size must be a positive packing-buffer size "
238- f"(got { shuffle_buffer_size } )"
203+ "shuffle_buffer_size must be a positive packing-buffer size " f"(got { shuffle_buffer_size } )"
239204 )
240205 # Consume the public `shuffle` argument for config compatibility, but
241206 # do not allocate DynamicCutSampler's second, reservoir-style buffer.
@@ -252,19 +217,13 @@ def __init__(
252217 self ._inject_restored_packing_buffer = False
253218
254219 def _uses_indexed_restore (self ) -> bool :
255- return bool (self .cuts ) and all (
256- getattr (source , "has_constant_time_access" , False ) for source in self .cuts
257- )
220+ return bool (self .cuts ) and all (getattr (source , "has_constant_time_access" , False ) for source in self .cuts )
258221
259222 @staticmethod
260223 def _capture_packing_buffer_tokens (buffer ) -> list [tuple [Any , ...]]:
261224 saved = []
262225 for example_or_tuple in buffer :
263- examples = (
264- example_or_tuple
265- if isinstance (example_or_tuple , tuple )
266- else (example_or_tuple ,)
267- )
226+ examples = example_or_tuple if isinstance (example_or_tuple , tuple ) else (example_or_tuple ,)
268227 tokens = tuple (get_graph_origin (example ) for example in examples )
269228 if any (token is None for token in tokens ):
270229 raise RuntimeError (
@@ -278,13 +237,9 @@ def state_dict(self) -> dict[str, Any]:
278237 state = super ().state_dict ()
279238 if self ._uses_indexed_restore ():
280239 if self ._batcher is not None :
281- state ["packing_buffer_tokens" ] = self ._capture_packing_buffer_tokens (
282- self ._batcher .reuse_cuts_buffer
283- )
240+ state ["packing_buffer_tokens" ] = self ._capture_packing_buffer_tokens (self ._batcher .reuse_cuts_buffer )
284241 else :
285- state ["packing_buffer_tokens" ] = list (
286- self ._restored_packing_buffer_tokens
287- )
242+ state ["packing_buffer_tokens" ] = list (self ._restored_packing_buffer_tokens )
288243 else :
289244 # Replay restoration deterministically rebuilds the post-filter pool.
290245 state ["packing_buffer_tokens" ] = None
@@ -329,25 +284,18 @@ def _restore_packing_buffer(self) -> list[tuple[Any, ...]]:
329284 f"{ len (tokens )} != { len (active_sources )} ."
330285 )
331286 restored .append (
332- tuple (
333- resolve_iterator_source (source )[token ]
334- for source , token in zip (active_sources , tokens )
335- )
287+ tuple (resolve_iterator_source (source )[token ] for source , token in zip (active_sources , tokens ))
336288 )
337289 restored .extend (self ._restored_legacy_examples )
338290 return restored
339291
340292 def _initialize_epoch_iterator (self , * , rebuild_sources : bool ) -> None :
341293 if rebuild_sources or self ._active_cuts is None :
342294 self ._active_cuts = self ._make_epoch_sources ()
343- source_iterators = [
344- iter (resolve_iterator_source (source )) for source in self ._active_cuts
345- ]
295+ source_iterators = [iter (resolve_iterator_source (source )) for source in self ._active_cuts ]
346296 filtered_examples = Filter (
347297 iterator = zip (* source_iterators ),
348- predicate = lambda examples : all (
349- self ._filter_fn (example ) for example in examples
350- ),
298+ predicate = lambda examples : all (self ._filter_fn (example ) for example in examples ),
351299 diagnostics = self .diagnostics ,
352300 )
353301 self ._batcher = ExactTokenBatcher (
0 commit comments