import os import sys import time import heapq import semver from multiprocessing import Pool KILO = 1024 KILOBYTE = 1 * KILO MEGABYTE = KILO * KILOBYTE GIGABYTE = KILO * MEGABYTE TERABYTE = KILO * GIGABYTE SIZE_PREFIX_LOOKUP = {'k': KILOBYTE, 'm': MEGABYTE, 'g': GIGABYTE, 't': TERABYTE} def parse_file_size(file_size): file_size = file_size.lower().strip() if len(file_size) == 0: return 0 n = int(keep_only_digits(file_size)) if file_size[-1] == 'b': file_size = file_size[:-1] e = file_size[-1] return SIZE_PREFIX_LOOKUP[e] * n if e in SIZE_PREFIX_LOOKUP else n def keep_only_digits(txt): return ''.join(filter(str.isdigit, txt)) def secs_to_hours(secs): hours, remainder = divmod(secs, 3600) minutes, seconds = divmod(remainder, 60) return '%d:%02d:%02d' % (hours, minutes, seconds) def check_ctcdecoder_version(): ds_version_s = open(os.path.join(os.path.dirname(__file__), '../VERSION')).read().strip() try: # pylint: disable=import-outside-toplevel from ds_ctcdecoder import __version__ as decoder_version except ImportError as e: if e.msg.find('__version__') > 0: print("DeepSpeech version ({ds_version}) requires CTC decoder to expose __version__. " "Please upgrade the ds_ctcdecoder package to version {ds_version}".format(ds_version=ds_version_s)) sys.exit(1) raise e decoder_version_s = decoder_version.decode() rv = semver.compare(ds_version_s, decoder_version_s) if rv != 0: print("DeepSpeech version ({}) and CTC decoder version ({}) do not match. " "Please ensure matching versions are in use.".format(ds_version_s, decoder_version_s)) sys.exit(1) return rv class Interleaved: """Collection that lazily combines sorted collections in an interleaving fashion. During iteration the next smallest element from all the sorted collections is always picked. The collections must support iter() and len().""" def __init__(self, *iterables, key=lambda obj: obj): self.iterables = iterables self.key = key self.len = sum(map(len, iterables)) def __iter__(self): return heapq.merge(*self.iterables, key=self.key) def __len__(self): return self.len class LimitingPool: """Limits unbound ahead-processing of multiprocessing.Pool's imap method before items get consumed by the iteration caller. This prevents OOM issues in situations where items represent larger memory allocations.""" def __init__(self, processes=None, process_ahead=None, sleeping_for=0.1): self.process_ahead = os.cpu_count() if process_ahead is None else process_ahead self.sleeping_for = sleeping_for self.processed = 0 self.pool = Pool(processes=processes) def __enter__(self): return self def _limit(self, it): for obj in it: while self.processed >= self.process_ahead: time.sleep(self.sleeping_for) self.processed += 1 yield obj def imap(self, fun, it): for obj in self.pool.imap(fun, self._limit(it)): self.processed -= 1 yield obj def __exit__(self, exc_type, exc_value, traceback): self.pool.close() class ExceptionBox: """Helper class for passing-back and re-raising an exception from inside a TensorFlow dataset generator. Used in conjunction with `remember_exception`.""" def __init__(self): self.exception = None def raise_if_set(self): if self.exception is not None: exception = self.exception self.exception = None raise exception # pylint: disable = raising-bad-type def remember_exception(iterable, exception_box=None): """Wraps a TensorFlow dataset generator for catching its actual exceptions that would otherwise just interrupt iteration w/o bubbling up.""" def do_iterate(): try: yield from iterable() except StopIteration: return except Exception as ex: # pylint: disable = broad-except exception_box.exception = ex return iterable if exception_box is None else do_iterate