Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 23 additions & 16 deletions proximal/experimental/optimize/absorb.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,22 +4,22 @@
from proximal.experimental.models import ProxFn


def absorbFFTConv(prox_fn: SumSquares) -> SumSquares | LeastSquaresFFT:
def absorbFFTConv(prox_fn: SumSquares) -> tuple[bool, SumSquares | LeastSquaresFFT]:
is_lin_ops_empty: bool = len(prox_fn.lin_ops) == 0
if is_lin_ops_empty:
# Nothing to absorb
return prox_fn
return False, prox_fn

lin_op = prox_fn.lin_ops[-1]
has_fft_conv: bool = isinstance(lin_op, FFTConv)
if not has_fft_conv:
# Only FFTConv is supported. Skipping...
return prox_fn
return False, prox_fn

# A trick to force the static analyzer to recognize the FFTConv type
assert isinstance(lin_op, FFTConv)

return LeastSquaresFFT(
return True, LeastSquaresFFT(
alpha=prox_fn.alpha,
gamma=prox_fn.gamma,
# todo: pre-compute FFT
Expand All @@ -29,34 +29,34 @@ def absorbFFTConv(prox_fn: SumSquares) -> SumSquares | LeastSquaresFFT:
)


def absorbMultiplyAdd(prox_fn: ProxFn) -> ProxFn:
def absorbMultiplyAdd(prox_fn: ProxFn) -> tuple[bool, ProxFn]:
"""Absorb (a * x + b) into the proximal function."""

if len(prox_fn.lin_ops) == 0 or not isinstance(prox_fn.lin_ops[-1], MultiplyAdd):
return prox_fn
return False, prox_fn

scale = prox_fn.lin_ops[-1].scale
offset = prox_fn.lin_ops[-1].offset
prox_fn.beta *= scale
prox_fn.b = prox_fn.b - prox_fn.beta * offset
prox_fn.lin_ops = prox_fn.lin_ops[:-1]

return prox_fn
return True, prox_fn


def absorbCrop(prox_fn: SumSquares) -> ProxFn:
def absorbCrop(prox_fn: SumSquares) -> tuple[bool, ProxFn]:
"""sum_square(Crop(u)) -> WeighteddLeastSquare(u)."""

if len(prox_fn.lin_ops) == 0 or not isinstance(prox_fn.lin_ops[-1], Crop):
return prox_fn
return False, prox_fn

# Generate the values of the binary mask representing the crop
crop_op: Crop = prox_fn.lin_ops[-1]

def mask(x: int, y: int) -> float:
return (crop_op.left <= x < crop_op.left + crop_op.width) and (crop_op.top <= x < crop_op.top + crop_op.height)

return WeightedLeastSquares(
return True, WeightedLeastSquares(
lin_ops=prox_fn.lin_ops[:-1],
alpha=prox_fn.alpha,
beta=prox_fn.beta,
Expand All @@ -69,13 +69,20 @@ def mask(x: int, y: int) -> float:
def absorb(problem: Problem) -> Problem:
assert problem.omega_fn is None, "Problem is already split, why?"

for i, psi_fn in enumerate(problem.psi_fns):
problem.psi_fns[i] = absorbMultiplyAdd(psi_fn)
for i in range(len(problem.psi_fns)):
should_retry = True
while should_retry:
psi_fn = problem.psi_fns[i]
should_retry, psi_fn = absorbMultiplyAdd(psi_fn)

if isinstance(psi_fn, SumSquares):
problem.psi_fns[i] = absorbFFTConv(psi_fn)
if isinstance(psi_fn, SumSquares):
is_success, psi_fn = absorbFFTConv(psi_fn)
should_retry = should_retry or is_success

if isinstance(psi_fn, SumSquares):
problem.psi_fns[i] = absorbCrop(psi_fn)
if isinstance(psi_fn, SumSquares):
is_success, psi_fn = absorbCrop(psi_fn)
should_retry = should_retry or is_success

problem.psi_fns[i] = psi_fn

return problem
18 changes: 10 additions & 8 deletions proximal/experimental/optimize/group.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,26 +11,28 @@ def group(problem: Problem) -> Problem:
assert problem.omega_fn is None

new_prox_fns: list[ProxFn] = []
absorbed_list: list[int] = []

for prox_fn in problem.psi_fns:
if isinstance(prox_fn, SumSquares):
new_prox_fns.append(prox_fn)
for i, prox_fn in enumerate(problem.psi_fns):
if i in absorbed_list:
continue

current_lin_ops = hash(prox_fn.lin_ops)
for j, prox_fn2 in enumerate(problem.psi_fns):
if i == j or j in absorbed_list:
continue

candidate_lin_ops = hash(prox_fn2.lin_ops)

is_equivalent_lin_ops: bool = candidate_lin_ops == current_lin_ops
is_sum_squares: bool = isinstance(prox_fn2, SumSquares)
is_zero_offset: bool = isinstance(prox_fn2.b, float) and prox_fn2.b == 0.0
scale_is_one: bool = prox_fn2.beta == 1.0

if (not is_equivalent_lin_ops) or (not is_sum_squares) or (not is_zero_offset) or (not scale_is_one):
continue

prox_fn.gamma = prox_fn2.alpha
problem.psi_fns.pop(j)
if is_equivalent_lin_ops and is_sum_squares and is_zero_offset and scale_is_one:
prox_fn.gamma = prox_fn2.alpha
absorbed_list.append(j)
break

new_prox_fns.append(prox_fn)

Expand Down