(self, urls, seed, nshards=sys.maxsize, deterministic=True)
| 54 | |
| 55 | class ConfiguredResampledShards(ResampledShards): |
| 56 | def __init__(self, urls, seed, nshards=sys.maxsize, deterministic=True): |
| 57 | from sat.helpers import print_rank0 |
| 58 | |
| 59 | try: |
| 60 | from megatron.core.parallel_state import get_data_parallel_group |
| 61 | |
| 62 | group = get_data_parallel_group() |
| 63 | print_rank0("Using megatron data parallel group.") |
| 64 | except: |
| 65 | from sat.mpu import get_data_parallel_group |
| 66 | |
| 67 | try: |
| 68 | group = get_data_parallel_group() |
| 69 | print_rank0("Using sat data parallel group.") |
| 70 | except AssertionError: |
| 71 | group = None |
| 72 | print_rank0("No data parallel group is specified!") |
| 73 | worker_seed_sat_this = partial(worker_seed_sat, group=group, seed=seed) |
| 74 | super().__init__(urls, nshards, worker_seed_sat_this, deterministic) |
| 75 | |
| 76 | |
| 77 | class SimpleDistributedWebDataset(DataPipeline): |
no test coverage detected