Skip to content

Commit 8dae7fa

Browse files
committed
feat: allow soft bound restraints
1 parent 98f0cfc commit 8dae7fa

3 files changed

Lines changed: 85 additions & 13 deletions

File tree

‎src/diffpy/apps/refinebase/refinement_server.py‎

Lines changed: 21 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -489,8 +489,8 @@ async def solve(
489489
"and the second dict is variable-constraint_equation pair."
490490
),
491491
] = None,
492-
restraints: Annotated[
493-
list[str], "List of restraints to apply during the refinement"
492+
bounds: Annotated[
493+
dict, "Dictionary of bounds for variables or equations"
494494
] = None,
495495
name: Annotated[str, "Name of the refinement session"] = None,
496496
weights: Annotated[
@@ -519,8 +519,24 @@ async def solve(
519519
constraints : list[dict], optional
520520
First dict is new_variable-initial value pair,
521521
and the second dict is variable-constraint_equation pair.
522-
restraints : list[str], optional
523-
List of restraints to apply during the refinement.
522+
bounds : dict, optional
523+
Dictionary of bounds for variables or equations.
524+
e.g. {"variable_name":
525+
{
526+
"lower_bound": 0,
527+
"upper_bound": 10,
528+
"uncertainty": 1,
529+
"scaled": False
530+
}}
531+
# start copied from diffpy.srfit docstring
532+
scaled : bool, optional
533+
If True, the restraint penalty is scaled by the unrestrained
534+
point-average chi^2 (chi^2/numpoints) (default is False).
535+
params : dict, optional
536+
The dictionary of Parameters, indexed by name, that are used in
537+
`param_or_eq` (if an equation string is used) but are not part
538+
of the RecipeOrganizer (default is {}).
539+
# end copied from diffpy.srfit docstring
524540
name : str, optional
525541
Name of the refinement session.
526542
weights : list[float], optional
@@ -547,7 +563,7 @@ async def solve(
547563
variable_names,
548564
residual_equations=residual_equations,
549565
constraints=constraints,
550-
restraints=restraints,
566+
bounds=bounds,
551567
name=name,
552568
weights=weights,
553569
metas=metas,

‎src/diffpy/apps/refinebase/refinement_session.py‎

Lines changed: 19 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -447,19 +447,19 @@ def _solve(
447447
models,
448448
variable_names,
449449
constraints=None,
450-
restraints=None,
450+
bounds=None,
451451
weights=None,
452452
residual_equations=None,
453453
metas=None,
454-
verbose_iterations=0,
455454
):
456-
# NOTE: restraints to be implemented
457455
recipe = FitRecipe()
458456
self.recipes_dict[name] = recipe
459457
if weights is None:
460458
weights = numpy.ones(len(profiles))
461459
if residual_equations is None:
462460
residual_equations = ["chiv"] * len(profiles)
461+
if bounds is None:
462+
bounds = {}
463463
if metas is not None:
464464
for i in range(len(metas)):
465465
profiles[i].meta.update(metas[i])
@@ -494,7 +494,20 @@ def _solve(
494494
if var in recipe._parameters.values():
495495
continue
496496
recipe.add_variable(var, name=variable_names[i])
497-
497+
for eq_or_var_name, arg_dict in bounds.items():
498+
lb = arg_dict.get("lower_bound", -numpy.inf)
499+
ub = arg_dict.get("upper_bound", numpy.inf)
500+
use_soft_bounds = arg_dict.get("use_soft_bounds", True)
501+
if use_soft_bounds:
502+
uncertainty = arg_dict.get("uncertainty", 1)
503+
scaled = arg_dict.get("scaled", False)
504+
eq_or_var_name = eq_or_var_name.replace(".", "_")
505+
recipe.add_soft_bounds(
506+
eq_or_var_name, lb, ub, sig=uncertainty, scaled=scaled
507+
)
508+
else:
509+
par = self.get_variable(eq_or_var_name)["obj"]
510+
par.bound_range(lb, ub)
498511
recipe.free("all")
499512
leastsq(recipe.residual, recipe.getValues())
500513
# NOTE: non-scalar value will raise error in `get_results_string`
@@ -511,12 +524,11 @@ def solve(
511524
variable_names=[],
512525
residual_equations=None,
513526
constraints=None,
514-
restraints=None,
527+
bounds=None,
515528
name=uuid.uuid4(),
516529
weights=None,
517530
metas=None,
518531
include_sgpars=False,
519-
verbose_iterations=0,
520532
):
521533
profiles = []
522534
for profile_name in profile_names:
@@ -555,11 +567,10 @@ def solve(
555567
variable_names=variable_names,
556568
residual_equations=residual_equations,
557569
constraints=constraints,
558-
restraints=restraints,
570+
bounds=bounds,
559571
name=name,
560572
weights=weights,
561573
metas=metas,
562-
verbose_iterations=verbose_iterations,
563574
)
564575

565576
def plot(self):

‎tests/test_refinement_session.py‎

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -13,6 +13,51 @@
1313
_DATA_DIR = Path(__file__).parent / "data"
1414

1515

16+
def test_bonds():
17+
# C1: Create a soft bound with an upper limit and high penalty
18+
# Expect the variable to be constrained by the upper bound
19+
xarray = numpy.linspace(0, 10, 100)
20+
yarray = 2 * xarray
21+
session = RefinementSession()
22+
session.add_profile_from_arrays(
23+
profile_name="linear", xarray=xarray, yarray=yarray
24+
)
25+
session.add_equation_model(model_name="linear_model", equation_str="m*x")
26+
session.set_variables_value(
27+
{
28+
"linear_model.m": 1,
29+
}
30+
)
31+
session._solve(
32+
name="linear_upper_bound",
33+
profiles=[session.profiles_dict["linear"]],
34+
models=[session.models_dict["linear_model"]],
35+
variable_names=["linear_model.m"],
36+
bounds={"linear_model.m": {"upper_bound": 1.8, "uncertainty": 1e-4}},
37+
)
38+
assert numpy.isclose(
39+
session.get_variable("linear_model.m")["value"],
40+
1.8,
41+
rtol=1e-2,
42+
)
43+
# C2: Create a soft bound with a lower limit and high penalty
44+
# Expect the variable to be constrained by the lower bound
45+
session._solve(
46+
name="linear_lower_bound",
47+
profiles=[session.profiles_dict["linear"]],
48+
models=[session.models_dict["linear_model"]],
49+
variable_names=["linear_model.m"],
50+
bounds={"linear_model.m": {"lower_bound": 2.2, "uncertainty": 1e-4}},
51+
)
52+
assert numpy.isclose(
53+
session.get_variable("linear_model.m")["value"],
54+
2.2,
55+
rtol=1e-2,
56+
)
57+
# NOTE: hard bounds are defined but not used by diffpy.srfit
58+
# Skip testing hard bounds.
59+
60+
1661
def test_refine_sine(sine_profile):
1762
# C1: Refinement session without additional calculator or functions
1863
session = RefinementSession()

0 commit comments

Comments
 (0)