Skip to content

Commit f450266

Browse files
authored
Check that fitted parameters have reasonable range (#206)
1 parent b6d4216 commit f450266

6 files changed

Lines changed: 62 additions & 37 deletions

File tree

ratapi/examples/normal_reflectivity/DSPC_standard_layers.ipynb

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -85,7 +85,7 @@
8585
" Parameter(name=\"SAM Heads Hydration\", min=10.0, value=45.45, max=50.0, fit=True, prior_type=\"gaussian\", mu=30.0, sigma=3.0),\n",
8686
" #\n",
8787
" Parameter(name=\"CW Thickness\", min=10.0, value=17.12, max=28.0, fit=True, prior_type=\"uniform\"),\n",
88-
" Parameter(name=\"CW SLD\", min=0.0, value=0.0, max=1e-09, fit=False, prior_type=\"uniform\"),\n",
88+
" Parameter(name=\"CW SLD\", min=0.0, value=0.0, max=0.0, fit=False, prior_type=\"uniform\"),\n",
8989
" Parameter(name=\"CW Hydration\", min=99.9, value=100.0, max=100.0, fit=False, prior_type=\"uniform\"),\n",
9090
" #\n",
9191
" Parameter(name=\"Bilayer Heads Thickness\", min=7.0, value=10.7, max=17.0, fit=True, prior_type=\"gaussian\", mu=10.0, sigma=2.0),\n",
@@ -175,8 +175,8 @@
175175
"outputs": [],
176176
"source": [
177177
"del problem.scalefactors[0]\n",
178-
"problem.scalefactors.append(name=\"Scalefactor 1\", min=0.05, value=0.10, max=0.2, fit=False)\n",
179-
"problem.scalefactors.append(name=\"Scalefactor 2\", min=0.05, value=0.15, max=0.2, fit=False)\n",
178+
"problem.scalefactors.append(name=\"Scalefactor 1\", min=0.05, value=0.10, max=0.2, fit=True)\n",
179+
"problem.scalefactors.append(name=\"Scalefactor 2\", min=0.05, value=0.15, max=0.2, fit=True)\n",
180180
"\n",
181181
"# Now deal with the backgrounds\n",
182182
"del problem.backgrounds[0]\n",

ratapi/examples/normal_reflectivity/DSPC_standard_layers.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ def DSPC_standard_layers():
2222
problem.parameters.append(name="SAM Tails Hydration", min=1.0, value=5.252, max=50.0, fit=True)
2323
problem.parameters.append(name="SAM Roughness", min=1.0, value=5.64, max=15.0, fit=True)
2424
problem.parameters.append(name="CW Thickness", min=10.0, value=17.12, max=28.0, fit=True)
25-
problem.parameters.append(name="CW SLD", min=0.0, value=0.0, max=1e-09, fit=False)
25+
problem.parameters.append(name="CW SLD", min=0.0, value=0.0, max=0.0, fit=False)
2626

2727
problem.parameters.append(
2828
name="SAM Heads Thickness",
@@ -131,8 +131,8 @@ def DSPC_standard_layers():
131131

132132
# Set the scalefactors - use one for each contrast
133133
del problem.scalefactors[0]
134-
problem.scalefactors.append(name="Scalefactor 1", min=0.05, value=0.10, max=0.2, fit=False)
135-
problem.scalefactors.append(name="Scalefactor 2", min=0.05, value=0.15, max=0.2, fit=False)
134+
problem.scalefactors.append(name="Scalefactor 1", min=0.05, value=0.10, max=0.2, fit=True)
135+
problem.scalefactors.append(name="Scalefactor 2", min=0.05, value=0.15, max=0.2, fit=True)
136136

137137
# Now deal with the backgrounds
138138
del problem.backgrounds[0]

ratapi/inputs.py

Lines changed: 36 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -3,14 +3,15 @@
33
import importlib
44
import os
55
import pathlib
6+
import warnings
67
from collections.abc import Callable
78

89
import numpy as np
910

1011
import ratapi
1112
import ratapi.wrappers
1213
from ratapi.rat_core import Checks, Control, NameStore, ProblemDefinition
13-
from ratapi.utils.enums import Calculations, Languages, LayerModels, TypeOptions
14+
from ratapi.utils.enums import Calculations, Languages, LayerModels, Procedures, TypeOptions
1415

1516
parameter_field = {
1617
"parameters": "params",
@@ -137,19 +138,21 @@ def make_input(project: ratapi.Project, controls: ratapi.Controls) -> tuple[Prob
137138
The controls object used in the compiled RAT code.
138139
139140
"""
140-
problem = make_problem(project)
141+
problem = make_problem(project, controls.procedure != Procedures.Calculate)
141142
cpp_controls = make_controls(controls)
142143

143144
return problem, cpp_controls
144145

145146

146-
def make_problem(project: ratapi.Project) -> ProblemDefinition:
147+
def make_problem(project: ratapi.Project, validate_range: bool = False) -> ProblemDefinition:
147148
"""Construct the problem input required for the compiled RAT code.
148149
149150
Parameters
150151
----------
151152
project : RAT.Project
152153
The project model, which defines the physical system under study.
154+
validate_range : bool, default True
155+
Whether parameter range should be validated.
153156
154157
Returns
155158
-------
@@ -351,40 +354,43 @@ def make_problem(project: ratapi.Project) -> ProblemDefinition:
351354
problem.domainContrastLayers = [
352355
domain_contrast_model if domain_contrast_model else [] for domain_contrast_model in domain_contrast_models
353356
]
354-
problem.fitParams = [
355-
param.value
356-
for class_list in ratapi.project.parameter_class_lists
357-
for param in getattr(project, class_list)
358-
if param.fit
359-
]
360-
problem.fitLimits = [
361-
[param.min, param.max]
362-
for class_list in ratapi.project.parameter_class_lists
363-
for param in getattr(project, class_list)
364-
if param.fit
365-
]
366-
problem.priorNames = [
367-
param.name for class_list in ratapi.project.parameter_class_lists for param in getattr(project, class_list)
368-
]
369-
problem.priorValues = [
370-
[prior_id[param.prior_type], param.mu, param.sigma]
371-
for class_list in ratapi.project.parameter_class_lists
372-
for param in getattr(project, class_list)
373-
]
357+
358+
fit_params = []
359+
fit_limits = []
360+
prior_names = []
361+
prior_values = []
362+
problem.checks = Checks()
363+
for class_list in ratapi.project.parameter_class_lists:
364+
field = parameter_field[class_list]
365+
check_list = []
366+
for param in getattr(project, class_list):
367+
prior_names.append(param.name)
368+
prior_values.append([prior_id[param.prior_type], param.mu, param.sigma])
369+
check_list.append(int(param.fit))
370+
if param.fit:
371+
min_range = 1e-6 if param.value == 0 else abs(param.value) * 1e-6
372+
if validate_range and (param.max - param.min) < min_range:
373+
warnings.warn(
374+
f'{class_list.replace("_", " ").title()} "{param.name}" was removed from the '
375+
f"fit because its range is too small (< {min_range:g}).",
376+
stacklevel=2,
377+
)
378+
check_list[-1] = 0
379+
else:
380+
fit_params.append(param.value)
381+
fit_limits.append([param.min, param.max])
382+
setattr(problem.checks, field, check_list)
383+
problem.fitParams = fit_params
384+
problem.fitLimits = fit_limits
385+
problem.priorNames = prior_names
386+
problem.priorValues = prior_values
374387

375388
# Names
376389
problem.names = NameStore()
377390
for class_list in ratapi.project.parameter_class_lists:
378391
setattr(problem.names, parameter_field[class_list], [param.name for param in getattr(project, class_list)])
379392
problem.names.contrasts = [contrast.name for contrast in project.contrasts]
380393

381-
# Checks
382-
problem.checks = Checks()
383-
for class_list in ratapi.project.parameter_class_lists:
384-
setattr(
385-
problem.checks, parameter_field[class_list], [int(element.fit) for element in getattr(project, class_list)]
386-
)
387-
388394
check_indices(problem)
389395

390396
return problem

ratapi/run.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -130,7 +130,9 @@ def run(project, controls):
130130
# Update parameter values in project
131131
for class_list in ratapi.project.parameter_class_lists:
132132
for index, value in enumerate(getattr(problem_definition, parameter_field[class_list])):
133-
getattr(project, class_list)[index].value = value
133+
param = getattr(project, class_list)[index]
134+
param.fit = bool(getattr(problem_definition.checks, parameter_field[class_list])[index])
135+
param.value = value
134136

135137
controls.delete_IPC()
136138

tests/test_data/R1DSPCBilayer.mat

-18 Bytes
Binary file not shown.

tests/test_inputs.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -483,6 +483,23 @@ def test_make_problem(test_project, test_problem, request) -> None:
483483
check_problem_equal(problem, test_problem)
484484

485485

486+
@pytest.mark.parametrize("value,min_range", [(0, 1e-6), (10, 1e-5)])
487+
def test_make_problem_validate_range(value, min_range, request) -> None:
488+
"""The problem should not contain fitted parameters with small range."""
489+
test_project = request.getfixturevalue("standard_layers_project")
490+
491+
test_project.scalefactors.set_fields(0, min=value, value=value, max=value, fit=True)
492+
problem = make_problem(test_project)
493+
assert problem.checks.scalefactors[0] == 1
494+
495+
with pytest.warns(
496+
UserWarning,
497+
match="was removed from the fit because its range is too small {0}(< {1:g}{0})".format("\\", min_range),
498+
):
499+
problem = make_problem(test_project, True)
500+
assert problem.checks.scalefactors[0] == 0
501+
502+
486503
@pytest.mark.parametrize("test_problem", ["standard_layers_problem", "custom_xy_problem", "domains_problem"])
487504
class TestCheckIndices:
488505
"""Tests for check_indices over a set of three test problems."""

0 commit comments

Comments
 (0)