MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / ConfiguredResampledShards

Class ConfiguredResampledShards

SwissArmyTransformer/sat/data_utils/webds.py:56–72  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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

Callers 3

__init__Method · 0.70
__init__Method · 0.70
__init__Method · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected