| 30 | |
| 31 | |
| 32 | class AudioIterableDataset(IterableDataset): |
| 33 | def __init__(self, file_list, flag): |
| 34 | self.file_list = file_list |
| 35 | self.flag = flag |
| 36 | |
| 37 | def parse_file_list(self): |
| 38 | worker_info = torch.utils.data.get_worker_info() |
| 39 | with open(self.file_list, "r") as f: |
| 40 | parsed_files = [{"src": line.strip()} for line in f] |
| 41 | |
| 42 | if worker_info: |
| 43 | # split workload |
| 44 | worker_id = worker_info.id |
| 45 | num_workers = worker_info.num_workers |
| 46 | parsed_files = parsed_files[worker_id::num_workers] |
| 47 | |
| 48 | return parsed_files |
| 49 | |
| 50 | def url_opener(self, data): |
| 51 | for sample in data: |
| 52 | assert "src" in sample |
| 53 | url = sample["src"] |
| 54 | try: |
| 55 | pr = urlparse(url) |
| 56 | # local file |
| 57 | if pr.scheme == "" or pr.scheme == "file": |
| 58 | stream = open(url, "rb") |
| 59 | # network file, such as HTTP(HDFS/OSS/S3)/HTTPS/SCP |
| 60 | else: |
| 61 | cmd = f"wget -q -O - {url}" |
| 62 | process = Popen(cmd, shell=True, stdout=PIPE) |
| 63 | sample.update(process=process) |
| 64 | stream = process.stdout |
| 65 | sample.update(stream=stream) |
| 66 | yield sample |
| 67 | except Exception as ex: |
| 68 | logging.warning("Failed to open {}".format(url)) |
| 69 | |
| 70 | def tar_file_and_group(self, sample): |
| 71 | assert "stream" in sample |
| 72 | stream = None |
| 73 | results = [] |
| 74 | try: |
| 75 | stream = tarfile.open(fileobj=sample["stream"], mode="r:*") |
| 76 | prev_prefix = None |
| 77 | example = {} |
| 78 | valid = True |
| 79 | for tarinfo in stream: |
| 80 | name = tarinfo.name |
| 81 | pos = name.rfind(".") |
| 82 | assert pos > 0 |
| 83 | prefix, postfix = name[:pos], name[pos + 1 :] |
| 84 | if prev_prefix is not None and prefix != prev_prefix: |
| 85 | example["key"] = prev_prefix |
| 86 | if valid: |
| 87 | results.append(example) |
| 88 | example = {} |
| 89 | valid = True |