MCPcopy Create free account
hub / github.com/LuChengTHU/dpm-solver / parallel_data_prefetch

Function parallel_data_prefetch

examples/stable-diffusion/ldm/util.py:108–203  ·  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

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

Callers 2

load_datapoolFunction · 0.90
load_databaseMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected