Skip to content

Commit 4344ab7

Browse files
committed
added option to turn off parallelization
1 parent 52db843 commit 4344ab7

3 files changed

Lines changed: 31 additions & 17 deletions

File tree

pyproject.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@ build-backend = "hatchling.build"
44

55
[project]
66
name = "pyGroupedTransforms"
7-
version = "0.1.0"
7+
version = "0.2.0"
88
authors = [
99
{ name="Felix Wirth", email="fwi012001@gmail.com" },
1010
]

src/pyGroupedTransforms/GroupedTransform.py

Lines changed: 27 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -184,6 +184,7 @@ def __init__(
184184
X,
185185
settings=[],
186186
fastmult=True,
187+
parallel=True,
187188
basis_vect=[],
188189
N=[],
189190
U=None,
@@ -232,6 +233,7 @@ def __init__(
232233
self.basis_vect = basis_vect
233234
self.system = system
234235
self.X = X
236+
self.parallel = parallel
235237

236238
if len(settings) == 0:
237239
self.settings = get_setting(system=system, N=N, U=U, d=d, ds=ds)
@@ -322,20 +324,25 @@ def __mul__(self, other):
322324

323325
if isinstance(other, np.ndarray): # `f = F*f` (f = other)
324326
if self.fastmult:
325-
threads = []
326327
fhat = GroupedCoefficients(self.settings)
327328

328329
def adjoint_worker(i):
329330
adjoint_result = self.transforms[i].H @ other
330331
fhat[self.settings[i].u] = adjoint_result
331-
332-
for i in range(len(self.transforms)):
333-
t = threading.Thread(target=adjoint_worker, args=(i,))
334-
t.start()
335-
threads.append(t)
336-
337-
for t in threads:
338-
t.join()
332+
333+
if self.parallel:
334+
threads = []
335+
for i in range(len(self.transforms)):
336+
t = threading.Thread(target=adjoint_worker, args=(i,))
337+
t.start()
338+
threads.append(t)
339+
340+
for t in threads:
341+
t.join()
342+
343+
else:
344+
for i in range(len(self.transforms)):
345+
adjoint_worker(i)
339346

340347
return fhat
341348
else:
@@ -349,20 +356,24 @@ def adjoint_worker(i):
349356
)
350357

351358
if self.fastmult:
352-
threads = []
353359
results = []
354360

355361
def worker(i):
356362
u = self.settings[i].u
357363
result = self.transforms[i] @ other[u]
358364
results.append(result)
359365

360-
for i in range(len(self.transforms)):
361-
t = threading.Thread(target=worker, args=(i,))
362-
t.start()
363-
threads.append(t)
364-
for t in threads:
365-
t.join()
366+
if self.parallel:
367+
threads = []
368+
for i in range(len(self.transforms)):
369+
t = threading.Thread(target=worker, args=(i,))
370+
t.start()
371+
threads.append(t)
372+
for t in threads:
373+
t.join()
374+
else:
375+
for i in range(len(self.transforms)):
376+
worker(i)
366377

367378
return sum(results)
368379
else:

src/pyGroupedTransforms/__init__.py

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,4 +35,7 @@
3535
"Setting",
3636
"GroupedCoefficientsComplex",
3737
"GroupedCoefficientsReal",
38+
"get_superposition_set", #TODO: Exports ab hier testen
39+
"GroupedTransform",
40+
"GroupedCoefficients",
3841
]

0 commit comments

Comments
 (0)