MCPcopy Create free account
hub / github.com/bbaaii/DreamDiffusion / parallel_data_prefetch

Function parallel_data_prefetch

code/dc_ldm/util.py:107–202  ·  view source on GitHub ↗
(
        func: callable, data, n_proc, target_data_type="ndarray", cpu_intensive=True, use_worker_id=False
)

Source from the content-addressed store, hash-verified

105
106
107def parallel_data_prefetch(
108 func: callable, data, n_proc, target_data_type="ndarray", cpu_intensive=True, use_worker_id=False
109):
110 # if target_data_type not in ["ndarray", "list"]:
111 # raise ValueError(
112 # "Data, which is passed to parallel_data_prefetch has to be either of type list or ndarray."
113 # )
114 if isinstance(data, np.ndarray) and target_data_type == "list":
115 raise ValueError("list expected but function got ndarray.")
116 elif isinstance(data, abc.Iterable):
117 if isinstance(data, dict):
118 print(
119 f'WARNING:"data" argument passed to parallel_data_prefetch is a dict: Using only its values and disregarding keys.'
120 )
121 data = list(data.values())
122 if target_data_type == "ndarray":
123 data = np.asarray(data)
124 else:
125 data = list(data)
126 else:
127 raise TypeError(
128 f"The data, that shall be processed parallel has to be either an np.ndarray or an Iterable, but is actually {type(data)}."
129 )
130
131 if cpu_intensive:
132 Q = mp.Queue(1000)
133 proc = mp.Process
134 else:
135 Q = Queue(1000)
136 proc = Thread
137 # spawn processes
138 if target_data_type == "ndarray":
139 arguments = [
140 [func, Q, part, i, use_worker_id]
141 for i, part in enumerate(np.array_split(data, n_proc))
142 ]
143 else:
144 step = (
145 int(len(data) / n_proc + 1)
146 if len(data) % n_proc != 0
147 else int(len(data) / n_proc)
148 )
149 arguments = [
150 [func, Q, part, i, use_worker_id]
151 for i, part in enumerate(
152 [data[i: i + step] for i in range(0, len(data), step)]
153 )
154 ]
155 processes = []
156 for i in range(n_proc):
157 p = proc(target=_do_parallel_data_prefetch, args=arguments[i])
158 processes += [p]
159
160 # start processes
161 print(f"Start prefetching...")
162 import time
163
164 start = time.time()

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected